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

Method pack_sequence

data/dataset_base.py:390–636  ·  view source on GitHub ↗
(self, sample, sequence_status)

Source from the content-addressed store, hash-verified

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(

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