#!/usr/bin/env python3 """ Landvex Vision — Träningspipeline Kontrastivt lärande från Zoomer-foton och tillverkarbilder. """ import json from pathlib import Path from typing import List, Dict, Tuple from dataclasses import dataclass import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader from torchvision import models, transforms from PIL import Image import numpy as np from tqdm import tqdm @dataclass class Traeningsexempel: bild_path: Path lvx_id: str positiv: bool # True = matchar lvx_id, False = negativt exempel kalla: str # "zoomer", "tillverkare", "syntetisk" class LandvexDataset(Dataset): """Dataset för kontrastivt lärande.""" def __init__(self, exempel: List[Traeningsexempel], transform=None): self.exempel = exempel self.transform = transform or transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), ]) def __len__(self): return len(self.exempel) def __getitem__(self, idx): ex = self.exempel[idx] bild = Image.open(ex.bild_path).convert('RGB') if self.transform: bild = self.transform(bild) return bild, ex.lvx_id, ex.positiv, ex.kalla class KontrastivtVerlust(nn.Module): """NT-Xent loss (Normalized Temperature-scaled Cross Entropy).""" def __init__(self, temperatur: float = 0.5): super().__init__() self.temperatur = temperatur self.cos_sim = nn.CosineSimilarity(dim=-1) def forward(self, z_i: torch.Tensor, z_j: torch.Tensor) -> torch.Tensor: """ z_i, z_j: normaliserade embeddings [batch_size, dim] """ # Cosine similarity sim = self.cos_sim(z_i.unsqueeze(1), z_j.unsqueeze(0)) / self.temperatur # Positiva par är på diagonalen etiketter = torch.arange(len(z_i)).to(z_i.device) # Cross entropy return nn.functional.cross_entropy(sim, etiketter) class LandvexTrainer: """Träningspipeline för Identify-modellen.""" def __init__( self, model: nn.Module, device: torch.device, learning_rate: float = 1e-4, temperatur: float = 0.5 ): self.model = model.to(device) self.device = device self.optimizer = optim.Adam(model.parameters(), lr=learning_rate) self.criterion = KontrastivtVerlust(temperatur) self.epoch = 0 def trana_epok(self, dataloader: DataLoader) -> dict: """Träna en epok.""" self.model.train() total_loss = 0 antal_batch = 0 for bilder, lvx_ids, positiva, kallor in tqdm(dataloader, desc=f"Epok {self.epoch}"): bilder = bilder.to(self.device) # Forward pass embeddings = self.model(bilder) # Kontrastivt förlust # Dela i två vyer (augmentation) batch_size = len(bilder) // 2 z_i = embeddings[:batch_size] z_j = embeddings[batch_size:] loss = self.criterion(z_i, z_j) # Backward pass self.optimizer.zero_grad() loss.backward() self.optimizer.step() total_loss += loss.item() antal_batch += 1 self.epoch += 1 return { "epok": self.epoch, "medel_loss": total_loss / antal_batch, } def utvardera(self, dataloader: DataLoader) -> dict: """Utvärdera modellen.""" self.model.eval() korrekta = 0 totala = 0 with torch.no_grad(): for bilder, lvx_ids, positiva, kallor in dataloader: bilder = bilder.to(self.device) embeddings = self.model(bilder) # TODO: Implementera top-k utvärdering totala += len(bilder) return { "noggrannhet": korrekta / totala if totala > 0 else 0, "antal": totala, } def spara(self, path: Path): """Spara träningsstatus.""" path.parent.mkdir(parents=True, exist_ok=True) torch.save({ "epoch": self.epoch, "model_state": self.model.state_dict(), "optimizer_state": self.optimizer.state_dict(), }, path) print(f" 💾 Träningsstatus sparad: {path}") def ladda(self, path: Path): """Ladda träningsstatus.""" checkpoint = torch.load(path, map_location=self.device) self.model.load_state_dict(checkpoint["model_state"]) self.optimizer.load_state_dict(checkpoint["optimizer_state"]) self.epoch = checkpoint["epoch"] print(f" 📂 Träningsstatus laddad: epok {self.epoch}") def skapa_syntetiskt_dataset(output_dir: Path, antal_klasser: int = 10, antal_bilder_per_klass: int = 20): """Skapa syntetiskt dataset för testning.""" from PIL import ImageDraw output_dir.mkdir(parents=True, exist_ok=True) exempel = [] for klass_idx in range(antal_klasser): lvx_id = f"LVX-TRP-{1000 + klass_idx:04d}" klass_dir = output_dir / lvx_id klass_dir.mkdir(exist_ok=True) for bild_idx in range(antal_bilder_per_klass): # Skapa syntetisk bild img = Image.new('RGB', (224, 224), color=(klass_idx * 20, 100, 150)) draw = ImageDraw.Draw(img) draw.rectangle([50, 50, 174, 174], outline="white", width=3) # Spara bild_path = klass_dir / f"{bild_idx:03d}.jpg" img.save(bild_path) exempel.append(Traeningsexempel( bild_path=bild_path, lvx_id=lvx_id, positiv=True, kalla="syntetisk" )) # Spara metadata with open(output_dir / "dataset.json", 'w') as f: json.dump([{ "bild_path": str(e.bild_path), "lvx_id": e.lvx_id, "positiv": e.positiv, "kalla": e.kalla, } for e in exempel], f, indent=2) return exempel def main(): """Demo: träna på syntetiskt dataset.""" print("🚀 Landvex Vision Trainer") # Skapa dataset dataset_dir = Path("/tmp/landvex-vision-dataset") exempel = skapa_syntetiskt_dataset(dataset_dir, antal_klasser=5, antal_bilder_per_klass=10) print(f"📊 Dataset: {len(exempel)} exempel") # Skapa modell device = torch.device("cuda" if torch.cuda.is_available() else "cpu") backbone = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2) backbone.fc = nn.Identity() # Ta bort klassificeringslager # Dataset och dataloader dataset = LandvexDataset(exempel) dataloader = DataLoader(dataset, batch_size=8, shuffle=True) # Tränare trainer = LandvexTrainer(backbone, device) # Träna for epok in range(3): resultat = trainer.trana_epok(dataloader) print(f" Epok {resultat['epok']}: loss = {resultat['medel_loss']:.4f}") # Spara trainer.spara(Path("/tmp/landvex-vision-checkpoint.pt")) print("\n✅ Träning klar!") if __name__ == "__main__": main()