File size: 2,717 Bytes
8cedc06
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
import pandas as pd
from torch.utils.data import DataLoader
from transformers import AutoTokenizer

from src.models import ImageEncoder, AudioEncoder, TextEncoder
from src.dataset import TriModalDataset, collate_fn
from src.utils import tri_modal_loss


DEVICE = "cuda" if torch.cuda.is_available() else "cpu"

EPOCHS = 10
BATCH_SIZE = 8
LR = 1e-4


def main():

    # Load metadata
    meta = pd.read_csv("ESC-50-master/meta/esc50.csv")

    meta = pd.read_csv("ESC-50-master/meta/esc50.csv")

    classes = sorted(meta["category"].unique())

    mini = meta[meta["category"].isin(classes)].copy()
    mini["text"] = mini["category"].str.replace("_", " ")

    print(f"Dataset size: {len(mini)}")

    tokenizer = AutoTokenizer.from_pretrained(
        "distilbert-base-uncased"
    )

    dataset = TriModalDataset(mini)

    loader = DataLoader(
        dataset,
        batch_size=BATCH_SIZE,
        shuffle=True,
        collate_fn=collate_fn(tokenizer)
    )

    img_encoder = ImageEncoder().to(DEVICE)
    aud_encoder = AudioEncoder().to(DEVICE)
    txt_encoder = TextEncoder().to(DEVICE)

    optimizer = torch.optim.AdamW(
        list(img_encoder.parameters()) +
        list(aud_encoder.parameters()) +
        list(txt_encoder.parameters()),
        lr=LR
    )

    print("Starting training...")

    for epoch in range(EPOCHS):

        img_encoder.train()
        aud_encoder.train()
        txt_encoder.train()

        total_loss = 0

        for imgs, auds, toks in loader:

            imgs = imgs.to(DEVICE)
            auds = auds.to(DEVICE)

            toks = {
                "input_ids": toks["input_ids"].to(DEVICE),
                "attention_mask": toks["attention_mask"].to(DEVICE)
            }

            optimizer.zero_grad()

            img_emb = img_encoder(imgs)
            aud_emb = aud_encoder(auds)
            txt_emb = txt_encoder(**toks)

            loss = tri_modal_loss(
                img_emb,
                aud_emb,
                txt_emb
            )

            loss.backward()
            optimizer.step()

            total_loss += loss.item()

        avg_loss = total_loss / len(loader)

        print(
            f"Epoch {epoch+1}/{EPOCHS} | Loss: {avg_loss:.4f}"
        )

    torch.save(
        {
            "image_encoder": img_encoder.state_dict(),
            "audio_encoder": aud_encoder.state_dict(),
            "text_encoder": txt_encoder.state_dict()
        },
        "trimodal_bind.pt"
    )

    print("Training complete.")
    print("Model saved as trimodal_bind.pt")


if __name__ == "__main__":
    main()