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

Method __init__

data/dataset_base.py:53–122  ·  view source on GitHub ↗
(
        self, 
        data_config, 
        tokenizer, 
        special_tokens,
        local_rank, 
        world_size, 
        num_workers,
        expected_num_tokens=32768, 
        max_num_tokens_per_sample=16384,
        max_num_tokens=36864,
        prefer_buffer_before=16384,
        max_buffer_size=50,
        interpolate_pos=False,
        use_flex=False,
        data_status=None,
    )

Source from the content-addressed store, hash-verified

51
52class PackedDataset(torch.utils.data.IterableDataset):
53 def __init__(
54 self,
55 data_config,
56 tokenizer,
57 special_tokens,
58 local_rank,
59 world_size,
60 num_workers,
61 expected_num_tokens=32768,
62 max_num_tokens_per_sample=16384,
63 max_num_tokens=36864,
64 prefer_buffer_before=16384,
65 max_buffer_size=50,
66 interpolate_pos=False,
67 use_flex=False,
68 data_status=None,
69 ):
70 super().__init__()
71 self.expected_num_tokens = expected_num_tokens
72 self.max_num_tokens_per_sample = max_num_tokens_per_sample
73 self.prefer_buffer_before = prefer_buffer_before
74 self.max_num_tokens = max_num_tokens
75 self.max_buffer_size = max_buffer_size
76 self.tokenizer = tokenizer
77 self.local_rank = local_rank
78 self.world_size = world_size
79 self.num_workers = num_workers
80 self.use_flex = use_flex
81 self.step_counter = 0
82 for k, v in special_tokens.items():
83 setattr(self, k, v)
84
85 #added for aug
86 self.cojitter = True #common_config.augs.cojitter
87 # Probability of using shared jitter vs. frame-specific jitter
88 self.cojitter_ratio = 0.3 #common_config.augs.cojitter_ratio
89 # Initialize image augmentations (color jitter, grayscale, gaussian blur)
90 self.image_aug = get_image_augmentation(
91 gray_scale=True,
92 gau_blur=False,
93 color_jitter=None,
94 )
95
96 self.resnet_normalize = transforms.Normalize(mean=_RESNET_MEAN, std=_RESNET_STD)
97
98 grouped_names, grouped_datasets, is_mandatory, grouped_weights = self.build_datasets(
99 data_config.grouped_datasets, data_status
100 )
101 self.grouped_datasets = grouped_datasets
102 self.dataset_iters = [(iter(dataset), grouped_name, dataset) for (dataset, grouped_name) in zip(grouped_datasets, grouped_names)]
103 self.is_mandatory = is_mandatory
104 self.grouped_weights = grouped_weights
105 self.data_config = data_config
106 self.interpolate_pos = interpolate_pos
107 if self.interpolate_pos:
108 self.get_flattened_position_ids = get_flattened_position_ids_interpolate
109 else:
110 self.get_flattened_position_ids = get_flattened_position_ids_extrapolate

Callers

nothing calls this directly

Calls 3

build_datasetsMethod · 0.95
get_image_augmentationFunction · 0.85
__init__Method · 0.45

Tested by

no test coverage detected