(metadata_dataset, map_fn, i, cache_dir, regenerate_cache, caching_batch_size)
| 175 | |
| 176 | |
| 177 | def _cache_text_embeddings(metadata_dataset, map_fn, i, cache_dir, regenerate_cache, caching_batch_size): |
| 178 | |
| 179 | def flatten_captions(example): |
| 180 | result = {key: [] for key in example} |
| 181 | for i, captions in enumerate(example['caption']): |
| 182 | for caption in captions: |
| 183 | result['caption'].append(caption) |
| 184 | for key, value in example.items(): |
| 185 | if key == 'caption': |
| 186 | continue |
| 187 | result[key].append(value[i]) |
| 188 | return result |
| 189 | |
| 190 | flattened_captions = metadata_dataset.map(flatten_captions, batched=True, keep_in_memory=True, remove_columns=metadata_dataset.column_names) |
| 191 | te_dataset = _map_and_cache( |
| 192 | flattened_captions, |
| 193 | map_fn, |
| 194 | cache_dir, |
| 195 | cache_file_prefix=f'text_embeddings_{i}_', |
| 196 | new_fingerprint_args=[i], |
| 197 | regenerate_cache=regenerate_cache, |
| 198 | caching_batch_size=caching_batch_size, |
| 199 | ) |
| 200 | assert len(te_dataset) == len(flattened_captions) |
| 201 | return TextEmbeddingDataset(te_dataset, flattened_captions) |
| 202 | |
| 203 | |
| 204 | # The smallest unit of a dataset. Represents a single size bucket from a single folder of images |
no test coverage detected