(self,
cfg,
phase='train',
collate_fn=all_collate,
batch_size: int = 32,
num_workers: int = 16,
**kwargs)
| 11 | class 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) |
nothing calls this directly
no test coverage detected