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