From 5fb29df06ccfe2dd36efe0d8944f1ca81c9cf0e7 Mon Sep 17 00:00:00 2001 From: Lars Bogner Date: Tue, 25 Nov 2025 11:06:10 +0100 Subject: [PATCH] Limit batch size for long context --- Code/python/notebooks/hyperparameter_scan.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/Code/python/notebooks/hyperparameter_scan.py b/Code/python/notebooks/hyperparameter_scan.py index f10578c..8d6afb3 100644 --- a/Code/python/notebooks/hyperparameter_scan.py +++ b/Code/python/notebooks/hyperparameter_scan.py @@ -70,14 +70,14 @@ def split_dataset(dataset, frac=0.8, seed=42): ) -def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epochs=100, altitude_weight=1e-3): +def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epochs=100, altitude_weight=1e-3, batch_size=256): optimizer = torch.optim.Adam(model.parameters(), lr=1e-2) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=5, factor=0.5) criterion = losses.HaversineAltitudeLoss(alt_const=altitude_weight) for epoch in range(epochs): batch_losses = [] - loader = torch.utils.data.DataLoader(dataset, batch_size=256, shuffle=True, collate_fn=dataloader.collate_to_cpu, num_workers=4) + loader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=True, collate_fn=dataloader.collate_to_cpu, num_workers=4) for X_f, X_t, y, X_c in loader: X_f = X_f.to(DEVICE) @@ -100,7 +100,7 @@ def train_one_model(model, dataset, val_dataset, warm_up_steps, pred_steps, epoc # Validation with torch.no_grad(): val_losses = [] - val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=256, collate_fn=dataloader.collate_to_cpu, num_workers=4) + val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=batch_size, collate_fn=dataloader.collate_to_cpu, num_workers=4) for X_f, X_t, y, X_c in val_loader: X_f = X_f.to(DEVICE) X_t = X_t.to(DEVICE) @@ -270,7 +270,7 @@ def run_mode_3(): print(f"Training {base_name} warm={warm} pred={pred} altitude_weight={altitude_weight}") - loss = train_one_model(model, train_ds, val_ds, warm // step, pred // step, altitude_weight=altitude_weight) + loss = train_one_model(model, train_ds, val_ds, warm // step, pred // step, altitude_weight=altitude_weight, batch_size=64) torch.save(model.state_dict(), f"{base_name}_wu{warm}_ps{pred}_aw{altitude_weight}.pt") results.append((base_name, warm, pred, altitude_weight, loss))