MCPcopy Create free account
hub / github.com/Zhiyuan-R/Tiger-Diffusion / get_dataset

Function get_dataset

train_generation.py:492–511  ·  view source on GitHub ↗
(dataroot, npoints,category)

Source from the content-addressed store, hash-verified

490
491
492def get_dataset(dataroot, npoints,category):
493 tr_dataset = ShapeNet15kPointClouds(root_dir=dataroot,
494 categories=[category], split='train',
495 tr_sample_size=npoints,
496 te_sample_size=npoints,
497 scale=1.,
498 normalize_per_shape=False,
499 normalize_std_per_axis=False,
500 random_subsample=True)
501 te_dataset = ShapeNet15kPointClouds(root_dir=dataroot,
502 categories=[category], split='val',
503 tr_sample_size=npoints,
504 te_sample_size=npoints,
505 scale=1.,
506 normalize_per_shape=False,
507 normalize_std_per_axis=False,
508 all_points_mean=tr_dataset.all_points_mean,
509 all_points_std=tr_dataset.all_points_std,
510 )
511 return tr_dataset, te_dataset
512
513
514def get_dataloader(opt, train_dataset, test_dataset=None):

Callers 2

trainFunction · 0.70
mainFunction · 0.70

Calls 1

Tested by

no test coverage detected