(
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,
)
| 51 | |
| 52 | class 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 |
nothing calls this directly
no test coverage detected