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}")