MCPcopy Create free account
hub / github.com/InternRobotics/G2VLM / pack_sequence

Method pack_sequence

data/dataset_base_periter.py:386–632  ·  view source on GitHub ↗
(self, sample, sequence_status)

Source from the content-addressed store, hash-verified

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(

Callers 1

__iter__Method · 0.95

Calls 6

len2weightFunction · 0.85
get_rope_index_image_3DFunction · 0.85
patchifyFunction · 0.85
getMethod · 0.80

Tested by

no test coverage detected