(image_list, data_list, save_folder)
| 114 | |
| 115 | |
| 116 | def create(image_list, data_list, save_folder): |
| 117 | assert image_list is not None, "image_list must be provided to generate features" |
| 118 | embed_size=512 |
| 119 | seg_maps = [] |
| 120 | total_lengths = [] |
| 121 | timer = 0 |
| 122 | img_embeds = torch.zeros((len(image_list), 300, embed_size)) |
| 123 | seg_maps = torch.zeros((len(image_list), 4, *image_list[0].shape[1:])) |
| 124 | mask_generator.predictor.model.to('cuda') |
| 125 | |
| 126 | for i, img in tqdm(enumerate(image_list), desc="Embedding images", leave=False): |
| 127 | timer += 1 |
| 128 | try: |
| 129 | img_embed, seg_map = _embed_clip_sam_tiles(img.unsqueeze(0), sam_encoder) |
| 130 | except: |
| 131 | raise ValueError(timer) |
| 132 | |
| 133 | lengths = [len(v) for k, v in img_embed.items()] |
| 134 | total_length = sum(lengths) |
| 135 | total_lengths.append(total_length) |
| 136 | |
| 137 | if total_length > img_embeds.shape[1]: |
| 138 | pad = total_length - img_embeds.shape[1] |
| 139 | img_embeds = torch.cat([ |
| 140 | img_embeds, |
| 141 | torch.zeros((len(image_list), pad, embed_size)) |
| 142 | ], dim=1) |
| 143 | |
| 144 | img_embed = torch.cat([v for k, v in img_embed.items()], dim=0) |
| 145 | assert img_embed.shape[0] == total_length |
| 146 | img_embeds[i, :total_length] = img_embed |
| 147 | |
| 148 | seg_map_tensor = [] |
| 149 | lengths_cumsum = lengths.copy() |
| 150 | for j in range(1, len(lengths)): |
| 151 | lengths_cumsum[j] += lengths_cumsum[j-1] |
| 152 | for j, (k, v) in enumerate(seg_map.items()): |
| 153 | if j == 0: |
| 154 | seg_map_tensor.append(torch.from_numpy(v)) |
| 155 | continue |
| 156 | assert v.max() == lengths[j] - 1, f"{j}, {v.max()}, {lengths[j]-1}" |
| 157 | v[v != -1] += lengths_cumsum[j-1] |
| 158 | seg_map_tensor.append(torch.from_numpy(v)) |
| 159 | seg_map = torch.stack(seg_map_tensor, dim=0) |
| 160 | seg_maps[i] = seg_map |
| 161 | |
| 162 | # mask_generator.predictor.model.to('cpu') |
| 163 | |
| 164 | # for i in range(img_embeds.shape[0]): |
| 165 | save_path = os.path.join(save_folder, data_list[i].split('.')[0]) |
| 166 | # assert total_lengths[i] == int(seg_maps[i].max() + 1) |
| 167 | curr = { |
| 168 | 'feature': img_embeds[i, :total_lengths[i]], |
| 169 | 'seg_maps': seg_maps[i] |
| 170 | } |
| 171 | sava_numpy(save_path, curr) |
| 172 | |
| 173 | def sava_numpy(save_path, data): |
no test coverage detected