MCPcopy Create free account
hub / github.com/SpatialVLA/SpatialVLA / get_train_dataloader

Function get_train_dataloader

train/monkey_patch.py:182–216  ·  view source on GitHub ↗

Returns the training [`~torch.utils.data.DataLoader`]. Will use no sampler if `train_dataset` does not implement `__len__`, a random sampler (adapted to distributed training if necessary) otherwise. Subclass and override this method if you want to inject some custom behavior.

(self)

Source from the content-addressed store, hash-verified

180 print('Replace train sampler!!')
181
182def get_train_dataloader(self) -> DataLoader:
183 """
184 Returns the training [`~torch.utils.data.DataLoader`].
185
186 Will use no sampler if `train_dataset` does not implement `__len__`, a random sampler (adapted to distributed
187 training if necessary) otherwise.
188
189 Subclass and override this method if you want to inject some custom behavior.
190 """
191 if self.train_dataset is None:
192 raise ValueError("Trainer: training requires a train_dataset.")
193
194 train_dataset = self.train_dataset
195 data_collator = self.data_collator
196 if is_datasets_available() and isinstance(train_dataset, datasets.Dataset):
197 train_dataset = self._remove_unused_columns(train_dataset, description="training")
198 else:
199 data_collator = self._get_collator_with_removed_columns(data_collator, description="training")
200
201 dataloader_params = {
202 "batch_size": self._train_batch_size,
203 "collate_fn": data_collator,
204 "num_workers": self.args.dataloader_num_workers,
205 "pin_memory": self.args.dataloader_pin_memory,
206 "persistent_workers": self.args.dataloader_persistent_workers,
207 }
208
209 if not isinstance(train_dataset, torch.utils.data.IterableDataset):
210 dataloader_params["sampler"] = self._get_train_sampler()
211 dataloader_params["drop_last"] = self.args.dataloader_drop_last
212 dataloader_params["worker_init_fn"] = seed_worker
213
214 if train_dataset.use_raw_dataloader:
215 return DataLoader(train_dataset, **dataloader_params)
216 return self.accelerator.prepare(DataLoader(train_dataset, **dataloader_params))
217
218def replace_train_dataloader():
219 transformers.Trainer.get_train_dataloader = get_train_dataloader

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected