Limit batch size for long context
This commit is contained in:
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user