fix bug in continual_learning for PyTorch* models
This commit is contained in:
@@ -74,16 +74,17 @@ class PyTorchMLPClassifier(BasePyTorchClassifier):
|
|||||||
model.to(self.device)
|
model.to(self.device)
|
||||||
optimizer = torch.optim.AdamW(model.parameters(), lr=self.learning_rate)
|
optimizer = torch.optim.AdamW(model.parameters(), lr=self.learning_rate)
|
||||||
criterion = torch.nn.CrossEntropyLoss()
|
criterion = torch.nn.CrossEntropyLoss()
|
||||||
init_model = self.get_init_model(dk.pair)
|
# check if continual_learning is activated, and retreive the model to continue training
|
||||||
trainer = PyTorchModelTrainer(
|
trainer = self.get_init_model(dk.pair)
|
||||||
model=model,
|
if trainer is None:
|
||||||
optimizer=optimizer,
|
trainer = PyTorchModelTrainer(
|
||||||
criterion=criterion,
|
model=model,
|
||||||
model_meta_data={"class_names": class_names},
|
optimizer=optimizer,
|
||||||
device=self.device,
|
criterion=criterion,
|
||||||
init_model=init_model,
|
model_meta_data={"class_names": class_names},
|
||||||
data_convertor=self.data_convertor,
|
device=self.device,
|
||||||
**self.trainer_kwargs,
|
data_convertor=self.data_convertor,
|
||||||
)
|
**self.trainer_kwargs,
|
||||||
|
)
|
||||||
trainer.fit(data_dictionary, self.splits)
|
trainer.fit(data_dictionary, self.splits)
|
||||||
return trainer
|
return trainer
|
||||||
|
|||||||
@@ -69,15 +69,16 @@ class PyTorchMLPRegressor(BasePyTorchRegressor):
|
|||||||
model.to(self.device)
|
model.to(self.device)
|
||||||
optimizer = torch.optim.AdamW(model.parameters(), lr=self.learning_rate)
|
optimizer = torch.optim.AdamW(model.parameters(), lr=self.learning_rate)
|
||||||
criterion = torch.nn.MSELoss()
|
criterion = torch.nn.MSELoss()
|
||||||
init_model = self.get_init_model(dk.pair)
|
# check if continual_learning is activated, and retreive the model to continue training
|
||||||
trainer = PyTorchModelTrainer(
|
trainer = self.get_init_model(dk.pair)
|
||||||
model=model,
|
if trainer is None:
|
||||||
optimizer=optimizer,
|
trainer = PyTorchModelTrainer(
|
||||||
criterion=criterion,
|
model=model,
|
||||||
device=self.device,
|
optimizer=optimizer,
|
||||||
init_model=init_model,
|
criterion=criterion,
|
||||||
data_convertor=self.data_convertor,
|
device=self.device,
|
||||||
**self.trainer_kwargs,
|
data_convertor=self.data_convertor,
|
||||||
)
|
**self.trainer_kwargs,
|
||||||
|
)
|
||||||
trainer.fit(data_dictionary, self.splits)
|
trainer.fit(data_dictionary, self.splits)
|
||||||
return trainer
|
return trainer
|
||||||
|
|||||||
@@ -75,17 +75,18 @@ class PyTorchTransformerRegressor(BasePyTorchRegressor):
|
|||||||
model.to(self.device)
|
model.to(self.device)
|
||||||
optimizer = torch.optim.AdamW(model.parameters(), lr=self.learning_rate)
|
optimizer = torch.optim.AdamW(model.parameters(), lr=self.learning_rate)
|
||||||
criterion = torch.nn.MSELoss()
|
criterion = torch.nn.MSELoss()
|
||||||
init_model = self.get_init_model(dk.pair)
|
# check if continual_learning is activated, and retreive the model to continue training
|
||||||
trainer = PyTorchTransformerTrainer(
|
trainer = self.get_init_model(dk.pair)
|
||||||
model=model,
|
if trainer is None:
|
||||||
optimizer=optimizer,
|
trainer = PyTorchTransformerTrainer(
|
||||||
criterion=criterion,
|
model=model,
|
||||||
device=self.device,
|
optimizer=optimizer,
|
||||||
init_model=init_model,
|
criterion=criterion,
|
||||||
data_convertor=self.data_convertor,
|
device=self.device,
|
||||||
window_size=self.window_size,
|
data_convertor=self.data_convertor,
|
||||||
**self.trainer_kwargs,
|
window_size=self.window_size,
|
||||||
)
|
**self.trainer_kwargs,
|
||||||
|
)
|
||||||
trainer.fit(data_dictionary, self.splits)
|
trainer.fit(data_dictionary, self.splits)
|
||||||
return trainer
|
return trainer
|
||||||
|
|
||||||
|
|||||||
@@ -25,7 +25,7 @@ class PyTorchModelTrainer(PyTorchTrainerInterface):
|
|||||||
optimizer: Optimizer,
|
optimizer: Optimizer,
|
||||||
criterion: nn.Module,
|
criterion: nn.Module,
|
||||||
device: str,
|
device: str,
|
||||||
init_model: Dict,
|
# init_model: Dict,
|
||||||
data_convertor: PyTorchDataConvertor,
|
data_convertor: PyTorchDataConvertor,
|
||||||
model_meta_data: Dict[str, Any] = {},
|
model_meta_data: Dict[str, Any] = {},
|
||||||
window_size: int = 1,
|
window_size: int = 1,
|
||||||
@@ -56,8 +56,8 @@ class PyTorchModelTrainer(PyTorchTrainerInterface):
|
|||||||
self.max_n_eval_batches: Optional[int] = kwargs.get("max_n_eval_batches", None)
|
self.max_n_eval_batches: Optional[int] = kwargs.get("max_n_eval_batches", None)
|
||||||
self.data_convertor = data_convertor
|
self.data_convertor = data_convertor
|
||||||
self.window_size: int = window_size
|
self.window_size: int = window_size
|
||||||
if init_model:
|
# if init_model:
|
||||||
self.load_from_checkpoint(init_model)
|
# self.load_from_checkpoint(init_model)
|
||||||
|
|
||||||
def fit(self, data_dictionary: Dict[str, pd.DataFrame], splits: List[str]):
|
def fit(self, data_dictionary: Dict[str, pd.DataFrame], splits: List[str]):
|
||||||
"""
|
"""
|
||||||
|
|||||||
Reference in New Issue
Block a user