(self, te_dataset, flattened_captions)
| 162 | |
| 163 | class TextEmbeddingDataset: |
| 164 | def __init__(self, te_dataset, flattened_captions): |
| 165 | self.te_dataset = te_dataset |
| 166 | self.flattened_captions = flattened_captions |
| 167 | self.image_spec_to_te_idx = defaultdict(list) |
| 168 | # TODO: maybe make this use Dataset object like the latents. But, you won't be caching text embeddings |
| 169 | # when training on very large datasets, so perhaps it doesn't really matter. |
| 170 | for i, image_spec in enumerate(flattened_captions['image_spec']): |
| 171 | self.image_spec_to_te_idx[tuple(image_spec)].append(i) |
| 172 | |
| 173 | def get_text_embeddings(self, image_spec, caption_number): |
| 174 | return self.te_dataset[self.image_spec_to_te_idx[image_spec][caption_number]] |
nothing calls this directly
no outgoing calls
no test coverage detected