pytorch - naming refactor - max_iters to n_steps

This commit is contained in:
yinon
2023-08-04 13:45:21 +00:00
parent d17bf6350d
commit a3c6904fbc
5 changed files with 12 additions and 12 deletions
@@ -26,7 +26,7 @@ class PyTorchMLPClassifier(BasePyTorchClassifier):
"model_training_parameters" : { "model_training_parameters" : {
"learning_rate": 3e-4, "learning_rate": 3e-4,
"trainer_kwargs": { "trainer_kwargs": {
"max_iters": 5000, "n_steps": 5000,
"batch_size": 64, "batch_size": 64,
"n_epochs": null, "n_epochs": null,
}, },
@@ -27,7 +27,7 @@ class PyTorchMLPRegressor(BasePyTorchRegressor):
"model_training_parameters" : { "model_training_parameters" : {
"learning_rate": 3e-4, "learning_rate": 3e-4,
"trainer_kwargs": { "trainer_kwargs": {
"max_iters": 5000, "n_steps": 5000,
"batch_size": 64, "batch_size": 64,
"n_epochs": null, "n_epochs": null,
}, },
@@ -30,7 +30,7 @@ class PyTorchTransformerRegressor(BasePyTorchRegressor):
"model_training_parameters" : { "model_training_parameters" : {
"learning_rate": 3e-4, "learning_rate": 3e-4,
"trainer_kwargs": { "trainer_kwargs": {
"max_iters": 5000, "n_steps": 5000,
"batch_size": 64, "batch_size": 64,
"n_epochs": null "n_epochs": null
}, },
@@ -39,7 +39,7 @@ class PyTorchModelTrainer(PyTorchTrainerInterface):
state_dict and model_meta_data saved by self.save() method. state_dict and model_meta_data saved by self.save() method.
:param model_meta_data: Additional metadata about the model (optional). :param model_meta_data: Additional metadata about the model (optional).
:param data_convertor: convertor from pd.DataFrame to torch.tensor. :param data_convertor: convertor from pd.DataFrame to torch.tensor.
:param max_iters: used to calculate n_epochs. The number of training iterations to run. :param n_steps: used to calculate n_epochs. The number of training iterations to run.
iteration here refers to the number of times optimizer.step() is called. iteration here refers to the number of times optimizer.step() is called.
ignored if n_epochs is set. ignored if n_epochs is set.
:param n_epochs: The maximum number batches to use for evaluation. :param n_epochs: The maximum number batches to use for evaluation.
@@ -50,10 +50,10 @@ class PyTorchModelTrainer(PyTorchTrainerInterface):
self.criterion = criterion self.criterion = criterion
self.model_meta_data = model_meta_data self.model_meta_data = model_meta_data
self.device = device self.device = device
self.max_iters: int = kwargs.get("max_iters", None) self.n_steps: int = kwargs.get("n_steps", None)
self.n_epochs: Optional[int] = kwargs.get("n_epochs", 10) self.n_epochs: Optional[int] = kwargs.get("n_epochs", 10)
if not self.max_iters and not self.n_epochs: if not self.n_steps and not self.n_epochs:
raise Exception("Either `max_iters` or `n_epochs` should be set.") raise Exception("Either `n_steps` or `n_epochs` should be set.")
self.batch_size: int = kwargs.get("batch_size", 64) self.batch_size: int = kwargs.get("batch_size", 64)
self.data_convertor = data_convertor self.data_convertor = data_convertor
@@ -82,7 +82,7 @@ class PyTorchModelTrainer(PyTorchTrainerInterface):
n_epochs = self.n_epochs or self.calc_n_epochs( n_epochs = self.n_epochs or self.calc_n_epochs(
n_obs=n_obs, n_obs=n_obs,
batch_size=self.batch_size, batch_size=self.batch_size,
n_iters=self.max_iters, n_iters=self.n_steps,
) )
batch_counter = 0 batch_counter = 0
@@ -153,7 +153,7 @@ class PyTorchModelTrainer(PyTorchTrainerInterface):
Calculates the number of epochs required to reach the maximum number Calculates the number of epochs required to reach the maximum number
of iterations specified in the model training parameters. of iterations specified in the model training parameters.
the motivation here is that `max_iters` is easier to optimize and keep stable, the motivation here is that `n_steps` is easier to optimize and keep stable,
across different n_obs - the number of data points. across different n_obs - the number of data points.
""" """
@@ -162,7 +162,7 @@ class PyTorchModelTrainer(PyTorchTrainerInterface):
if n_epochs <= 10: if n_epochs <= 10:
logger.warning( logger.warning(
f"Setting low n_epochs: {n_epochs}. " f"Setting low n_epochs: {n_epochs}. "
f"Please consider increasing `max_iters` hyper-parameter." f"Please consider increasing `n_steps` hyper-parameter."
) )
return n_epochs return n_epochs
+2 -2
View File
@@ -97,9 +97,9 @@ def mock_pytorch_mlp_model_training_parameters() -> Dict[str, Any]:
return { return {
"learning_rate": 3e-4, "learning_rate": 3e-4,
"trainer_kwargs": { "trainer_kwargs": {
"max_iters": 1, "n_steps": None,
"batch_size": 64, "batch_size": 64,
"n_epochs": None, "n_epochs": 1,
}, },
"model_kwargs": { "model_kwargs": {
"hidden_dim": 32, "hidden_dim": 32,