(self, sample, sequence_status)
| 388 | batch_data_indexes = [] |
| 389 | |
| 390 | def pack_sequence(self, sample, sequence_status): |
| 391 | if 'image_tensor_list' in sample: |
| 392 | image_tensor_list = sample['image_tensor_list'] |
| 393 | if 'image_grid_thw_list' in sample: |
| 394 | image_grid_thw_list = sample['image_grid_thw_list'] |
| 395 | text_ids_list = sample['text_ids_list'] |
| 396 | sequence_plan = sample['sequence_plan'] |
| 397 | |
| 398 | img_per_seq = 0 |
| 399 | if 'depths' in sample: |
| 400 | depth_array =sample['depths'] |
| 401 | if 'extrinsics' in sample: |
| 402 | extrinsics_array =sample['extrinsics'] |
| 403 | if 'intrinsics' in sample: |
| 404 | intrinsics_array =sample['intrinsics'] |
| 405 | if 'world_points' in sample: |
| 406 | world_points_array =sample['world_points'] |
| 407 | if 'point_masks' in sample: |
| 408 | point_masks_array =sample['point_masks'] |
| 409 | if 'view_infos' in sample: |
| 410 | view_infos =sample['view_infos'] |
| 411 | if 'image_paths' in sample: |
| 412 | image_paths =sample['image_paths'] |
| 413 | if 'img_per_seq' in sample: |
| 414 | img_per_seq = sample['img_per_seq'] |
| 415 | if 'dino_image_tensor_list' in sample: |
| 416 | dino_image_tensor_list = sample['dino_image_tensor_list'] |
| 417 | if 'dino_images' in sample: |
| 418 | dino_images = sample['dino_images'] |
| 419 | if 'dino_thw' in sample: |
| 420 | dino_thw = sample['dino_thw'] |
| 421 | |
| 422 | split_lens, attn_modes = list(), list() |
| 423 | curr = sequence_status['curr'] |
| 424 | curr_rope_id = 0 |
| 425 | sample_lens = 0 |
| 426 | vit_cnt = 0 |
| 427 | dino_cnt = 0 |
| 428 | |
| 429 | |
| 430 | for item in sequence_plan: |
| 431 | split_start = item.get('split_start', True) |
| 432 | if split_start: |
| 433 | curr_split_len = 0 |
| 434 | |
| 435 | if item['type'] == 'text': |
| 436 | text_ids = text_ids_list.pop(0) |
| 437 | if item['enable_cfg'] == 1 and random.random() < self.data_config.text_cond_dropout_prob: |
| 438 | continue |
| 439 | |
| 440 | shifted_text_ids = text_ids |
| 441 | sequence_status['packed_text_ids'].extend(shifted_text_ids) |
| 442 | sequence_status['packed_text_indexes'].extend(range(curr, curr + len(shifted_text_ids))) |
| 443 | |
| 444 | # \n text_token <-> text_token <|im_end|> |
| 445 | if item['loss'] == 1: |
| 446 | sequence_status['ce_loss_indexes'].extend(range(curr, curr + len(shifted_text_ids))) |
| 447 | sequence_status['ce_loss_weights'].extend( |
no test coverage detected