Limit batch size for long context

This commit is contained in:
2025-11-25 11:06:10 +01:00
parent 8098d6282d
commit 5fb29df06c
+4 -4
View File
@@ -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))