File size: 4,513 Bytes
d8c733f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 | from tqdm.notebook import tqdm
from Vathos.blocks import *
def symbolic_1d_ar_target_train(model, train_loader, val_loader, optimizer, scheduler, criterion,
device, NUM_EPOCHS=100, use_amp=False, CLIP_NORM=1.0, div=1):
scaler = torch.cuda.amp.GradScaler(enabled=use_amp)
steps_per_epoch = len(train_loader) // div
for epoch in range(NUM_EPOCHS):
model.train()
running_loss = 0.0
pbar = tqdm(train_loader, desc=f"Epoch {epoch + 1}/{NUM_EPOCHS} Train", total=steps_per_epoch)
steps = 0
for batch_x, batch_tgt in pbar:
batch_x = batch_x.to(device)
steps += 1
optimizer.zero_grad()
with torch.amp.autocast('cuda'):
logits = model(batch_x) # [B,S,V]
loss = criterion(
logits[:, :-1, :].reshape(-1, logits.size(-1)),
batch_x[:, 1:].reshape(-1)
)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_NORM)
scaler.step(optimizer)
scaler.update()
scheduler.step()
running_loss += loss.item()
lr_now = optimizer.param_groups[0]['lr']
pbar.set_postfix({'loss': running_loss / (pbar.n + 1), 'lr': lr_now})
if steps >= steps_per_epoch:
break
model.eval()
val_loss = 0.0
# with torch.no_grad():
if True:
for batch_x, batch_tgt, composer_idx in tqdm(val_loader, desc=f"Epoch {epoch + 1}/{NUM_EPOCHS} Val"):
batch_x = batch_x.to(device).requires_grad_(False)
logits = model(batch_x)
loss = criterion(
logits[:, :-1, :].reshape(-1, logits.size(-1)),
batch_x[:, 1:].reshape(-1)
)
val_loss += loss.item()
if isinstance(model, VathosModel):
model.module.register_loss(loss.item())
val_avg = val_loss / len(val_loader)
if isinstance(model, VathosModel):
model.module.register_epoch()
print(f"Epoch {epoch + 1} val avg loss: {val_avg:.4f}")
def symbolic_1d_ar_input_ids_train(model, train_loader, val_loader, optimizer, scheduler, criterion,
device, NUM_EPOCHS=100, use_amp=False, CLIP_NORM=1.0, spe=1000, val_steps=float('inf')):
scaler = torch.cuda.amp.GradScaler(enabled=use_amp)
steps_per_epoch = spe
for epoch in range(NUM_EPOCHS):
model.train()
running_loss = 0.0
pbar = tqdm(train_loader, desc=f"Epoch {epoch + 1}/{NUM_EPOCHS} Train", total=steps_per_epoch)
steps = 0
for batch_x in pbar:
batch_x = batch_x['input_ids'].to(device)
steps += 1
optimizer.zero_grad()
with torch.amp.autocast('cuda'):
logits = model(batch_x) # [B,S,V]
loss = criterion(
logits[:, :-1, :].reshape(-1, logits.size(-1)),
batch_x[:, 1:].reshape(-1)
)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), CLIP_NORM)
scaler.step(optimizer)
scaler.update()
if scheduler is not None:
scheduler.step()
running_loss += loss.item()
lr_now = optimizer.param_groups[0]['lr']
pbar.set_postfix({'loss': running_loss / (pbar.n + 1), 'lr': lr_now})
if steps >= steps_per_epoch:
break
model.eval()
val_loss = 0.0
with torch.no_grad():
# if True:
val_s = 0
for batch_x in tqdm(val_loader, desc=f"Epoch {epoch + 1}/{NUM_EPOCHS} Val"):
batch_x = batch_x['input_ids'].to(device).requires_grad_(False)
val_s += 1
if val_s >= val_steps:
break
logits = model(batch_x)
loss = criterion(
logits[:, :-1, :].reshape(-1, logits.size(-1)),
batch_x[:, 1:].reshape(-1)
)
model.module.register_loss(loss.item())
val_avg = model.module.get_mean_loss()
model.module.register_epoch()
print(f"Epoch {epoch + 1} val avg loss: {val_avg:.4f}")
|