(batch)
| 107 | |
| 108 | |
| 109 | def collate_fn_motion_wanvideo(batch): |
| 110 | # filter out None |
| 111 | batch = [x for x in batch if x is not None] |
| 112 | |
| 113 | # collate motion |
| 114 | motions = [x.pop('motion') for x in batch] |
| 115 | motions = collate_tensors(motions) |
| 116 | motion_lengths = [x.get('motion_length') for x in batch] |
| 117 | motion_mask = lengths_to_mask(motion_lengths, motions.device, motions.shape[1]) |
| 118 | |
| 119 | with_ref_motion = False |
| 120 | if 'ref_motion' in batch[0]: |
| 121 | with_ref_motion = True |
| 122 | ref_motions = [x.pop('ref_motion') for x in batch] |
| 123 | ref_motions_original = [x.pop('ref_motion_original') for x in batch] |
| 124 | ref_motions = collate_tensors(ref_motions) |
| 125 | ref_motions_original = collate_tensors(ref_motions_original) |
| 126 | ref_motion_lengths = [x.pop('ref_motion_length') for x in batch] |
| 127 | ref_motion_mask = lengths_to_mask(ref_motion_lengths, ref_motions.device, ref_motions.shape[1]) |
| 128 | |
| 129 | # collate prompt |
| 130 | with_prompt_emb = False |
| 131 | if 'prompt_emb' in batch[0]: |
| 132 | with_prompt_emb = True |
| 133 | prompt_emb = [x.pop('prompt_emb') for x in batch] |
| 134 | prompt_emb = collate_tensors(prompt_emb) |
| 135 | prompt_lengths = [x.pop('prompt_length') for x in batch] |
| 136 | prompt_emb_mask = lengths_to_mask(prompt_lengths, prompt_emb.device, prompt_emb.shape[1]) |
| 137 | |
| 138 | texts = None |
| 139 | if 'text' in batch[0]: |
| 140 | texts = [x.pop('text') for x in batch] |
| 141 | |
| 142 | sample_ids = None |
| 143 | if 'test_sample_id' in batch[0]: |
| 144 | sample_ids = [x.pop('test_sample_id') for x in batch] |
| 145 | |
| 146 | try: |
| 147 | ret = torch.utils.data.default_collate(batch) |
| 148 | except Exception as e: |
| 149 | print(f'Failed to collate batch in collate_fn_motion_wanvideo: {e}') |
| 150 | for idx, sample in enumerate(batch[:2]): |
| 151 | print(f'sample {idx} keys: {list(sample.keys())}') |
| 152 | raise |
| 153 | |
| 154 | ret['motion'] = motions |
| 155 | ret['motion_mask'] = motion_mask |
| 156 | if with_ref_motion: |
| 157 | ret['ref_motion'] = ref_motions |
| 158 | ret['ref_motion_mask'] = ref_motion_mask |
| 159 | ret['ref_motion_original'] = ref_motions_original |
| 160 | if with_prompt_emb: |
| 161 | ret['prompt_emb'] = prompt_emb |
| 162 | ret['prompt_emb_mask'] = prompt_emb_mask |
| 163 | if texts is not None: |
| 164 | ret['text'] = texts |
| 165 | if sample_ids is not None: |
| 166 | ret['test_sample_id'] = sample_ids |
nothing calls this directly
no test coverage detected