Fixed sampling for GRU and reduced batch size

This commit is contained in:
Victor Mylle
2024-01-19 00:10:12 +00:00
parent e8e53ab185
commit c6fa17fa40
2 changed files with 16 additions and 5 deletions

View File

@@ -38,10 +38,11 @@ data_config.NOMINAL_NET_POSITION = True
data_config = task.connect(data_config, name="data_features")
data_processor = DataProcessor(data_config, path="", lstm=True)
data_processor.set_batch_size(8192)
data_processor.set_batch_size(128)
data_processor.set_full_day_skip(True)
inputDim = data_processor.get_input_size()
print("Input dim: ", inputDim)
model_parameters = {
"epochs": 5000,
@@ -54,7 +55,7 @@ model_parameters = task.connect(model_parameters, name="model_parameters")
#### Model ####
# model = SimpleDiffusionModel(96, model_parameters["hidden_sizes"], other_inputs_dim=inputDim[1], time_dim=model_parameters["time_dim"])
model = GRUDiffusionModel(96, [256, 256], other_inputs_dim=inputDim[2], time_dim=64, gru_hidden_size=128)
model = GRUDiffusionModel(96, model_parameters["hidden_sizes"], other_inputs_dim=inputDim[2], time_dim=model_parameters["time_dim"], gru_hidden_size=256)
print("Starting training ...")