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)
| 180 | print('Replace train sampler!!') |
| 181 | |
| 182 | def 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 | |
| 218 | def replace_train_dataloader(): |
| 219 | transformers.Trainer.get_train_dataloader = get_train_dataloader |
nothing calls this directly
no outgoing calls
no test coverage detected