MCPcopy Create free account
hub / github.com/Xuehao-Gao/GUESS / __init__

Method __init__

mld/data/Kit.py:13–41  ·  view source on GitHub ↗
(self,
                 cfg,
                 phase='train',
                 collate_fn=all_collate,
                 batch_size: int = 32,
                 num_workers: int = 16,
                 **kwargs)

Source from the content-addressed store, hash-verified

11class KitDataModule(BASEDataModule):
12
13 def __init__(self,
14 cfg,
15 phase='train',
16 collate_fn=all_collate,
17 batch_size: int = 32,
18 num_workers: int = 16,
19 **kwargs):
20 super().__init__(batch_size=batch_size,
21 num_workers=num_workers,
22 collate_fn=collate_fn)
23 self.save_hyperparameters(logger=False)
24 self.name = 'kit'
25 self.njoints = 21
26 if phase == 'text_only':
27 self.Dataset = TextOnlyDataset
28 else:
29 self.Dataset = Text2MotionDatasetV2
30 self.cfg = cfg
31
32 sample_overrides = {
33 "split": "val",
34 "tiny": True,
35 "progress_bar": False
36 }
37 self._sample_set = self.get_sample_set(overrides=sample_overrides)
38
39 # Get additional info of the dataset
40 self.nfeats = self._sample_set.nfeats
41 # self.transforms = self._sample_set.transforms
42
43 def feats2joints(self, features):
44 mean = torch.tensor(self.hparams.mean).to(features)

Callers

nothing calls this directly

Calls 1

get_sample_setMethod · 0.80

Tested by

no test coverage detected