I'm currently using Huggingface's Trainer class to train Distillbert for a regression problem using a custom loss function. I'm using their checkpoints to resume training due to the ephemeral nature of compute / unexpected errors.
The issue I'm facing is that each time I resume training from a checkpoint as per their Trainer class via the model_path in the Trainer.train() method, I noticed that the class iterates over the dataloader until it reaches the iteration count as saved in the checkpoint (see the lines from the Trainer class that match the issue).
This might usually not be a issue, but due to the nature of my dataloader's collate function and the size of the dataset, iterating for such a duration without any training is pretty expensive and slows down the overall training.
I planned on utilizing a custom sampler class something along the lines of this with a parameter to resume the indices from a given location but that too seems quite the hack for the given problem.
What could be an alternative that I could try to save on this wasted compute cycles?