CoolFace
Apppublic

kalpkanungo/SceneGraphNet

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
train_relationship.py42 linesDownload Raw Back to training
1import torch2from torch.utils.data import DataLoader3from src.dataset import RelationshipDataset4from src.model import RelationshipNet5 6device = "mps" if torch.backends.mps.is_available() else "cpu"7 8dataset = RelationshipDataset(9    image_dir="data/relationship_dataset/images",10    label_path="data/relationship_dataset/labels_encoded.json"11)12 13loader = DataLoader(dataset, batch_size=16, shuffle=True)14 15num_classes = len(set([item["label"] for item in dataset.data]))16 17model = RelationshipNet(num_classes).to(device)18 19criterion = torch.nn.CrossEntropyLoss()20optimizer = torch.optim.Adam(model.parameters(), lr=3e-4)21 22epochs = 1023 24for epoch in range(epochs):25    total_loss = 026 27    for images, labels in loader:28        images = images.to(device)29        labels = labels.to(device)30 31        outputs = model(images)32        loss = criterion(outputs, labels)33 34        optimizer.zero_grad()35        loss.backward()36        optimizer.step()37 38        total_loss += loss.item()39 40    print(f"Epoch {epoch+1}, Loss: {total_loss:.4f}")41 42torch.save(model.state_dict(), "models/relationship_model.pth")