| 234 | return sequence_status |
| 235 | |
| 236 | def to_tensor(self, sequence_status): |
| 237 | data = dict( |
| 238 | sequence_length=sum(sequence_status['sample_lens']), |
| 239 | sample_lens=sequence_status['sample_lens'], |
| 240 | packed_text_ids=torch.tensor(sequence_status['packed_text_ids']), |
| 241 | packed_text_indexes=torch.tensor(sequence_status['packed_text_indexes']), |
| 242 | packed_position_ids=torch.cat(sequence_status['packed_position_ids'], dim=1), |
| 243 | ) |
| 244 | if not self.use_flex: |
| 245 | data['nested_attention_masks'] = sequence_status['nested_attention_masks'] |
| 246 | else: |
| 247 | sequence_len = data['sequence_length'] |
| 248 | pad_len = self.max_num_tokens - sequence_len #### this is fixed only postive num |
| 249 | data['split_lens'] = sequence_status['split_lens'] + [pad_len] |
| 250 | data['attn_modes'] = sequence_status['attn_modes'] + ['causal'] |
| 251 | data['sample_lens'] += [pad_len] |
| 252 | |
| 253 | if len(sequence_status['packed_dino_image_tensor_list']) > 0: |
| 254 | |
| 255 | data['packed_dino_token_indexes'] = torch.tensor(sequence_status['packed_dino_token_indexes']) |
| 256 | data['dino_token_seqlens'] = torch.tensor(sequence_status['dino_token_seqlens']) |
| 257 | |
| 258 | |
| 259 | packed_dino_image_tensors = torch.from_numpy(np.stack(sequence_status["packed_dino_image_tensor_list"]).astype(np.float32)).contiguous() |
| 260 | packed_dino_image_tensors = packed_dino_image_tensors.permute(0,3,1,2).to(torch.get_default_dtype()).div(255) |
| 261 | |
| 262 | if self.image_aug is not None: |
| 263 | if self.cojitter and random.random() > self.cojitter_ratio: |
| 264 | # Apply the same color jittering transformation to all frames |
| 265 | packed_dino_image_tensors = self.image_aug(packed_dino_image_tensors) |
| 266 | else: |
| 267 | # Apply different color jittering to each frame individually |
| 268 | for aug_img_idx in range(len(packed_dino_image_tensors)): |
| 269 | packed_dino_image_tensors[aug_img_idx] = self.image_aug(packed_dino_image_tensors[aug_img_idx]) |
| 270 | |
| 271 | packed_dino_image_tensors = packed_dino_image_tensors.contiguous() |
| 272 | |
| 273 | depths = torch.from_numpy(np.stack(sequence_status["packed_depths"]).astype(np.float32)).to(torch.float32) |
| 274 | extrinsics = torch.from_numpy(np.stack(sequence_status["packed_extrinsics"]).astype(np.float32)).to(torch.float32) |
| 275 | intrinsics = torch.from_numpy(np.stack(sequence_status["packed_intrinsics"]).astype(np.float32)).to(torch.float32) |
| 276 | world_points = torch.from_numpy(np.stack(sequence_status["packed_world_points"]).astype(np.float32)).to(torch.float32) |
| 277 | point_masks = torch.from_numpy(np.stack(sequence_status["packed_point_masks"])) # Mask indicating valid depths / world points / cam points per frame |
| 278 | |
| 279 | data["packed_depths"] = depths |
| 280 | data["packed_extrinsics"] = extrinsics |
| 281 | data["packed_intrinsics"] = intrinsics |
| 282 | # data["packed_cam_points"] = cam_points |
| 283 | data["packed_world_points"] = world_points |
| 284 | data["packed_point_masks"] = point_masks |
| 285 | # packed_dino_image_tensors = torch.stack(sequence_status["packed_dino_image_tensor_list"], dim=0) |
| 286 | |
| 287 | packed_dino_image_tensors = self.resnet_normalize(packed_dino_image_tensors) |
| 288 | data["packed_dino_image_tensor_list"] = packed_dino_image_tensors |
| 289 | data['img_per_seq_lens'] = sequence_status['img_per_seq_lens'] |
| 290 | data['packed_view_infos'] = sequence_status["packed_view_infos"] |
| 291 | data['packed_image_paths'] = sequence_status["packed_image_paths"] |
| 292 | |
| 293 | |