landvex: Datafabrik + Vision + Infrastruktur för 50GB skalning
Datafabrik: - Skördare: crawler, källvitlista, upphandlingsskördare - Extraktor: LLM-baserad schemastyrd extraktion - Upplösare: Entitetsupplösning och deduplicering - Köer: Schemalagd / kunddriven / fält - Agentorkestrering: 20+ parallella agenter Vision: - Identify-modell: ResNet50 + kontrastivt lärande - Träningspipeline: NT-Xent loss - Vektordatabas: FAISS för snabb sökning - OCR-pipeline: Typskyltsläsning Infrastruktur: - Docker Compose production - Terraform för AWS ECS - Prometheus + Grafana monitorering - Neo4j + FAISS + MinIO + Redis
This commit is contained in:
@@ -0,0 +1,32 @@
|
||||
FROM pytorch/pytorch:2.2.0-cuda12.1-cudnn8-runtime
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Systemberoenden
|
||||
RUN apt-get update && apt-get install -y \
|
||||
tesseract-ocr \
|
||||
tesseract-ocr-swe \
|
||||
tesseract-ocr-eng \
|
||||
libgl1-mesa-glx \
|
||||
libglib2.0-0 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Python-paket
|
||||
COPY requirements.txt .
|
||||
RUN pip install --no-cache-dir -r requirements.txt
|
||||
|
||||
# Kopiera kod
|
||||
COPY modeller/ ./modeller/
|
||||
COPY traening/ ./traening/
|
||||
COPY embeddings/ ./embeddings/
|
||||
COPY dataset/ ./dataset/
|
||||
|
||||
# Miljövariabler
|
||||
ENV PYTHONPATH=/app
|
||||
ENV VISION_MODEL_DIR=/models
|
||||
ENV VISION_DATASET_DIR=/data/dataset
|
||||
|
||||
VOLUME ["/models", "/data"]
|
||||
|
||||
# Default: träna
|
||||
CMD ["python3", "-m", "traening.trainer"]
|
||||
@@ -0,0 +1,179 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Landvex Vision — OCR-pipeline för typskyltar
|
||||
Extraherar tillverkare, modell, serienummer från foton.
|
||||
"""
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
|
||||
import pytesseract
|
||||
from PIL import Image
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
@dataclass
|
||||
class OCRResultat:
|
||||
text: str
|
||||
konfidens: float
|
||||
bbox: Tuple[int, int, int, int] # x, y, w, h
|
||||
falt: Dict[str, str] # Extraherade fält
|
||||
|
||||
class LandvexOCR:
|
||||
"""OCR-pipeline optimerad för infrastrukturtypskyltar."""
|
||||
|
||||
def __init__(self, tesseract_cmd: Optional[str] = None):
|
||||
if tesseract_cmd:
|
||||
pytesseract.pytesseract.tesseract_cmd = tesseract_cmd
|
||||
|
||||
# Fältmönster per objektklass
|
||||
self.falt_monster = {
|
||||
"LVX-ELN-0101": { # Kabelskåp
|
||||
"tillverkare": r"(ABB|Schneider|Siemens|GE)\s*[\w-]*",
|
||||
"modell": r"(CDC|Prisma|Okken|IM\s*\w+)",
|
||||
"ar": r"20\d{2}",
|
||||
},
|
||||
"LVX-TRP-0102": { # Vägbelysningsarmatur
|
||||
"tillverkare": r"(Philips|Thorn|Schreder|Acuity|Cree)",
|
||||
"modell": r"(SL-\w+|ER\w+|Vista\w+)",
|
||||
"effekt": r"(\d+)\s*W",
|
||||
},
|
||||
"LVX-VAT-0101": { # Brunnsbetäckning
|
||||
"tillverkare": r"(ULMA|ACO|Wrede|GDK)",
|
||||
"klass": r"(A15|B125|C250|D400|E600|F900)",
|
||||
"material": r"(gjutjärn|stål|komposit|betong)",
|
||||
},
|
||||
}
|
||||
|
||||
def forbehandla_bild(self, bild: np.ndarray) -> np.ndarray:
|
||||
"""Förbättra bildkvalitet för OCR."""
|
||||
# Konvertera till gråskala
|
||||
if len(bild.shape) == 3:
|
||||
gray = cv2.cvtColor(bild, cv2.COLOR_RGB2GRAY)
|
||||
else:
|
||||
gray = bild
|
||||
|
||||
# Brusreducering
|
||||
denoised = cv2.fastNlMeansDenoising(gray)
|
||||
|
||||
# Kontrastförbättring (CLAHE)
|
||||
clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))
|
||||
enhanced = clahe.apply(denoised)
|
||||
|
||||
# Skärpa
|
||||
kernel = np.array([[-1, -1, -1], [-1, 9, -1], [-1, -1, -1]])
|
||||
sharpened = cv2.filter2D(enhanced, -1, kernel)
|
||||
|
||||
return sharpened
|
||||
|
||||
def hitta_text_regioner(self, bild: np.ndarray) -> List[Tuple[int, int, int, int]]:
|
||||
"""Hitta regioner som troligen innehåller text."""
|
||||
# MSER (Maximally Stable Extremal Regions)
|
||||
mser = cv2.MSER_create()
|
||||
regions, _ = mser.detectRegions(bild)
|
||||
|
||||
# Filtrera små regioner
|
||||
h, w = bild.shape
|
||||
min_area = (h * w) * 0.001 # Minst 0.1% av bilden
|
||||
|
||||
bboxes = []
|
||||
for region in regions:
|
||||
x, y, w, h = cv2.boundingRect(region)
|
||||
if w * h > min_area and w > h * 2: # Text är oftast bredare än hög
|
||||
bboxes.append((x, y, w, h))
|
||||
|
||||
return bboxes
|
||||
|
||||
def kora_ocr(self, bild: Image.Image, lvx_id: Optional[str] = None) -> List[OCRResultat]:
|
||||
"""Kör OCR på bild och extrahera fält."""
|
||||
# Konvertera till numpy
|
||||
img_array = np.array(bild)
|
||||
|
||||
# Förbehandla
|
||||
processed = self.forbehandla_bild(img_array)
|
||||
|
||||
# Hitta textregioner
|
||||
regioner = self.hitta_text_regioner(processed)
|
||||
|
||||
resultat = []
|
||||
for x, y, w, h in regioner[:5]: # Max 5 regioner
|
||||
# Beskär region
|
||||
roi = processed[y:y+h, x:x+w]
|
||||
|
||||
# Kör Tesseract
|
||||
text = pytesseract.image_to_string(roi, lang='eng+swe')
|
||||
conf_data = pytesseract.image_to_data(roi, output_type=pytesseract.Output.DICT)
|
||||
|
||||
# Beräkna medelkonfidens
|
||||
konfidenser = [c for c in conf_data['conf'] if c > 0]
|
||||
medel_konf = np.mean(konfidenser) if konfidenser else 0
|
||||
|
||||
# Extrahera fält om vi vet objektklassen
|
||||
falt = {}
|
||||
if lvx_id and lvx_id in self.falt_monster:
|
||||
import re
|
||||
for falt_namn, monster in self.falt_monster[lvx_id].items():
|
||||
match = re.search(monster, text, re.IGNORECASE)
|
||||
if match:
|
||||
falt[falt_namn] = match.group(1)
|
||||
|
||||
resultat.append(OCRResultat(
|
||||
text=text.strip(),
|
||||
konfidens=medel_konf / 100.0, # Normalisera till 0-1
|
||||
bbox=(x, y, w, h),
|
||||
falt=falt
|
||||
))
|
||||
|
||||
# Sortera efter konfidens
|
||||
resultat.sort(key=lambda x: x.konfidens, reverse=True)
|
||||
return resultat
|
||||
|
||||
def extrahera_falt(self, text: str, lvx_id: str) -> Dict[str, str]:
|
||||
"""Extrahera strukturerade fält från OCR-text."""
|
||||
import re
|
||||
|
||||
falt = {}
|
||||
monster = self.falt_monster.get(lvx_id, {})
|
||||
|
||||
for falt_namn, pattern in monster.items():
|
||||
matches = re.findall(pattern, text, re.IGNORECASE)
|
||||
if matches:
|
||||
falt[falt_namn] = matches[0]
|
||||
|
||||
return falt
|
||||
|
||||
def main():
|
||||
"""Demo: OCR på syntetisk bild."""
|
||||
print("🔤 Landvex OCR Pipeline")
|
||||
|
||||
ocr = LandvexOCR()
|
||||
|
||||
# Skapa syntetisk testbild med text
|
||||
from PIL import ImageDraw, ImageFont
|
||||
img = Image.new('RGB', (400, 200), color='white')
|
||||
draw = ImageDraw.Draw(img)
|
||||
|
||||
try:
|
||||
font = ImageFont.truetype("/usr/share/fonts/truetype/dejavu/DejaVuSans-Bold.ttf", 24)
|
||||
except:
|
||||
font = ImageFont.load_default()
|
||||
|
||||
draw.text((20, 20), "ABB Kabeldon", fill='black', font=font)
|
||||
draw.text((20, 60), "CDC-LT 400", fill='black', font=font)
|
||||
draw.text((20, 100), "2023", fill='black', font=font)
|
||||
|
||||
# Kör OCR
|
||||
resultat = ocr.kora_ocr(img, lvx_id="LVX-ELN-0101")
|
||||
|
||||
print("\n🎯 OCR-resultat:")
|
||||
for i, r in enumerate(resultat[:3]):
|
||||
print(f" Region {i+1}:")
|
||||
print(f" Text: {r.text[:100]}")
|
||||
print(f" Konfidens: {r.konfidens:.2f}")
|
||||
print(f" Fält: {r.falt}")
|
||||
|
||||
print("\n✅ OCR klar!")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,140 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Landvex Vision — Vektordatabas
|
||||
Lagrar och söker bild- och textembeddings.
|
||||
"""
|
||||
import json
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Tuple, Optional
|
||||
from dataclasses import dataclass
|
||||
|
||||
import faiss # Facebook AI Similarity Search
|
||||
|
||||
@dataclass
|
||||
class VektorsokResultat:
|
||||
lvx_id: str
|
||||
distans: float
|
||||
metadata: dict
|
||||
|
||||
class LandvexVektordatabas:
|
||||
"""Vektordatabas för bild- och textembeddings."""
|
||||
|
||||
def __init__(self, dimension: int = 2048, index_type: str = "FlatIP"):
|
||||
self.dimension = dimension
|
||||
self.index_type = index_type
|
||||
|
||||
# FAISS index för snabb sökning
|
||||
if index_type == "FlatIP":
|
||||
self.index = faiss.IndexFlatIP(dimension) # Inner product (cosine om normaliserad)
|
||||
elif index_type == "IVF"::
|
||||
nlist = 100 # Antal kluster
|
||||
quantizer = faiss.IndexFlatIP(dimension)
|
||||
self.index = faiss.IndexIVFFlat(quantizer, dimension, nlist)
|
||||
else:
|
||||
raise ValueError(f"Okänd index-typ: {index_type}")
|
||||
|
||||
# Mappning index_id → lvx_id
|
||||
self.id_map: List[str] = []
|
||||
self.metadata: Dict[str, dict] = {}
|
||||
|
||||
def lagg_till(self, lvx_id: str, embedding: np.ndarray, metadata: dict):
|
||||
"""Lägg till embedding i databasen."""
|
||||
# Normalisera för cosine similarity
|
||||
embedding = embedding / np.linalg.norm(embedding)
|
||||
embedding = embedding.reshape(1, -1).astype('float32')
|
||||
|
||||
self.index.add(embedding)
|
||||
self.id_map.append(lvx_id)
|
||||
self.metadata[lvx_id] = metadata
|
||||
|
||||
def sok(self, query: np.ndarray, top_k: int = 5) -> List[VektorsokResultat]:
|
||||
"""Sök närmaste grannar."""
|
||||
query = query / np.linalg.norm(query)
|
||||
query = query.reshape(1, -1).astype('float32')
|
||||
|
||||
distanser, indices = self.index.search(query, top_k)
|
||||
|
||||
resultat = []
|
||||
for dist, idx in zip(distanser[0], indices[0]):
|
||||
if idx < 0 or idx >= len(self.id_map):
|
||||
continue
|
||||
|
||||
lvx_id = self.id_map[idx]
|
||||
resultat.append(VektorsokResultat(
|
||||
lvx_id=lvx_id,
|
||||
distans=float(dist),
|
||||
metadata=self.metadata.get(lvx_id, {})
|
||||
))
|
||||
|
||||
return resultat
|
||||
|
||||
def spara(self, path: Path):
|
||||
"""Spara index och metadata."""
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Spara FAISS-index
|
||||
faiss.write_index(self.index, str(path / "index.faiss"))
|
||||
|
||||
# Spara metadata
|
||||
with open(path / "metadata.json", 'w') as f:
|
||||
json.dump({
|
||||
"id_map": self.id_map,
|
||||
"metadata": self.metadata,
|
||||
"dimension": self.dimension,
|
||||
"index_type": self.index_type,
|
||||
}, f, indent=2)
|
||||
|
||||
print(f" 💾 Vektordatabas sparad: {path}")
|
||||
|
||||
def ladda(self, path: Path):
|
||||
"""Ladda index och metadata."""
|
||||
self.index = faiss.read_index(str(path / "index.faiss"))
|
||||
|
||||
with open(path / "metadata.json", 'r') as f:
|
||||
data = json.load(f)
|
||||
|
||||
self.id_map = data["id_map"]
|
||||
self.metadata = data["metadata"]
|
||||
self.dimension = data["dimension"]
|
||||
self.index_type = data["index_type"]
|
||||
|
||||
print(f" 📂 Vektordatabas laddad: {len(self.id_map)} vektorer")
|
||||
|
||||
def main():
|
||||
"""Demo: vektordatabas."""
|
||||
print("🔍 Landvex Vektordatabas")
|
||||
|
||||
db = LandvexVektordatabas(dimension=128) # Låg dimension för demo
|
||||
|
||||
# Lägg till exempel-embeddings
|
||||
np.random.seed(42)
|
||||
for i in range(100):
|
||||
emb = np.random.randn(128)
|
||||
db.lagg_till(
|
||||
lvx_id=f"LVX-TRP-{1000+i:04d}",
|
||||
embedding=emb,
|
||||
metadata={
|
||||
"namn_sv": f"Testobjekt {i}",
|
||||
"verifieringsniva": "kallbelagd",
|
||||
}
|
||||
)
|
||||
|
||||
# Sök
|
||||
query = np.random.randn(128)
|
||||
resultat = db.sok(query, top_k=5)
|
||||
|
||||
print("\n🎯 Sökresultat:")
|
||||
for r in resultat:
|
||||
print(f" {r.lvx_id}: distans={r.distans:.3f}, {r.metadata['namn_sv']}")
|
||||
|
||||
# Spara och ladda
|
||||
db.spara(Path("/tmp/landvex-vectordb"))
|
||||
|
||||
db2 = LandvexVektordatabas()
|
||||
db2.ladda(Path("/tmp/landvex-vectordb"))
|
||||
|
||||
print(f"\n✅ Databas: {len(db2.id_map)} vektorer")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,178 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Landvex Vision — Identify-modell
|
||||
Foto in → rankade kandidater med konfidens
|
||||
"""
|
||||
import json
|
||||
import hashlib
|
||||
from pathlib import Path
|
||||
from typing import List, Dict, Tuple, Optional
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchvision import models, transforms
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
@dataclass
|
||||
class IdentifyKandidat:
|
||||
lvx_id: str
|
||||
posttyp: str
|
||||
namn_sv: str
|
||||
namn_en: str
|
||||
konfidens: float
|
||||
verifieringsniva: str
|
||||
kannetecken_match: List[str]
|
||||
embedding_distans: float
|
||||
|
||||
class LandvexIdentifyModel:
|
||||
"""Vision-modell för infrastrukturidentifiering."""
|
||||
|
||||
def __init__(self, model_path: Optional[Path] = None):
|
||||
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
print(f" 🖥️ Enhet: {self.device}")
|
||||
|
||||
# Ladda förtränad ResNet som backbone
|
||||
self.backbone = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)
|
||||
self.backbone = nn.Sequential(*list(self.backbone.children())[:-1]) # Ta bort sista FC
|
||||
self.backbone = self.backbone.to(self.device)
|
||||
self.backbone.eval()
|
||||
|
||||
# Transform för bilder
|
||||
self.transform = 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]),
|
||||
])
|
||||
|
||||
# Embedding-databas (lvx_id → embedding)
|
||||
self.embeddings: Dict[str, np.ndarray] = {}
|
||||
self.metadata: Dict[str, dict] = {}
|
||||
|
||||
if model_path and model_path.exists():
|
||||
self.ladda(model_path)
|
||||
|
||||
def bild_till_embedding(self, bild: Image.Image) -> np.ndarray:
|
||||
"""Konvertera bild till embedding-vektor."""
|
||||
tensor = self.transform(bild).unsqueeze(0).to(self.device)
|
||||
|
||||
with torch.no_grad():
|
||||
embedding = self.backbone(tensor)
|
||||
|
||||
return embedding.cpu().numpy().flatten()
|
||||
|
||||
def lagg_till_klass(self, lvx_id: str, bilder: List[Image.Image], metadata: dict):
|
||||
"""Lägg till en objektklass med träningsbilder."""
|
||||
embeddings = [self.bild_till_embedding(b) for b in bilder]
|
||||
medel_embedding = np.mean(embeddings, axis=0)
|
||||
|
||||
self.embeddings[lvx_id] = medel_embedding / np.linalg.norm(medel_embedding)
|
||||
self.metadata[lvx_id] = metadata
|
||||
|
||||
def identifiera(self, bild: Image.Image, top_k: int = 5) -> List[IdentifyKandidat]:
|
||||
"""Identifiera objekt i bild. Returnera top-k kandidater."""
|
||||
if not self.embeddings:
|
||||
return []
|
||||
|
||||
query_embedding = self.bild_till_embedding(bild)
|
||||
query_embedding = query_embedding / np.linalg.norm(query_embedding)
|
||||
|
||||
# Beräkna kosinuslikhet
|
||||
resultat = []
|
||||
for lvx_id, emb in self.embeddings.items():
|
||||
distans = np.dot(query_embedding, emb)
|
||||
meta = self.metadata[lvx_id]
|
||||
|
||||
# Konfidens = likhet * verifieringsnivå-faktor
|
||||
verif_faktor = {
|
||||
"obekraftad": 0.5,
|
||||
"kallbelagd": 0.75,
|
||||
"faltverifierad": 0.9,
|
||||
"tillverkarbekraftad": 0.95,
|
||||
}.get(meta.get("verifieringsniva", "obekraftad"), 0.5)
|
||||
|
||||
konfidens = float(distans * verif_faktor)
|
||||
|
||||
resultat.append(IdentifyKandidat(
|
||||
lvx_id=lvx_id,
|
||||
posttyp=meta.get("posttyp", "objektklass"),
|
||||
namn_sv=meta.get("namn_sv", ""),
|
||||
namn_en=meta.get("namn_en", ""),
|
||||
konfidens=konfidens,
|
||||
verifieringsniva=meta.get("verifieringsniva", "obekraftad"),
|
||||
kannetecken_match=meta.get("kannetecken", [])[:3],
|
||||
embedding_distans=float(distans),
|
||||
))
|
||||
|
||||
# Sortera efter konfidens
|
||||
resultat.sort(key=lambda x: x.konfidens, reverse=True)
|
||||
return resultat[:top_k]
|
||||
|
||||
def spara(self, path: Path):
|
||||
"""Spara modell och embeddings."""
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
data = {
|
||||
"embeddings": {k: v.tolist() for k, v in self.embeddings.items()},
|
||||
"metadata": self.metadata,
|
||||
}
|
||||
torch.save(data, path)
|
||||
print(f" 💾 Modell sparad: {path}")
|
||||
|
||||
def ladda(self, path: Path):
|
||||
"""Ladda modell och embeddings."""
|
||||
data = torch.load(path, map_location=self.device)
|
||||
self.embeddings = {k: np.array(v) for k, v in data["embeddings"].items()}
|
||||
self.metadata = data["metadata"]
|
||||
print(f" 📂 Modell laddad: {len(self.embeddings)} klasser")
|
||||
|
||||
def main():
|
||||
"""Demo: träna på syntetiska data och identifiera."""
|
||||
from PIL import ImageDraw
|
||||
|
||||
model = LandvexIdentifyModel()
|
||||
|
||||
# Skapa syntetiska träningsbilder (i verkligheten: riktiga foton)
|
||||
def skapa_testbild(farg, storlek=(224, 224)):
|
||||
img = Image.new('RGB', storlek, farg)
|
||||
draw = ImageDraw.Draw(img)
|
||||
draw.rectangle([50, 50, 174, 174], outline="white", width=3)
|
||||
return img
|
||||
|
||||
# Lägg till två klasser
|
||||
model.lagg_till_klass("LVX-TRP-0101", [
|
||||
skapa_testbild("gray"),
|
||||
skapa_testbild("lightgray"),
|
||||
], {
|
||||
"posttyp": "objektklass",
|
||||
"namn_sv": "Belysningsstolpe",
|
||||
"namn_en": "Lighting column",
|
||||
"verifieringsniva": "kallbelagd",
|
||||
"kannetecken": ["Grå stolpe", "Ljustopp"],
|
||||
})
|
||||
|
||||
model.lagg_till_klass("LVX-TRP-0102", [
|
||||
skapa_testbild("blue"),
|
||||
skapa_testbild("darkblue"),
|
||||
], {
|
||||
"posttyp": "objektklass",
|
||||
"namn_sv": "Vägbelysningsarmatur",
|
||||
"namn_en": "Road lighting luminaire",
|
||||
"verifieringsniva": "kallbelagd",
|
||||
"kannetecken": ["Blå armatur", "LED-ljus"],
|
||||
})
|
||||
|
||||
# Testa identifiering
|
||||
test_bild = skapa_testbild("gray")
|
||||
kandidater = model.identifiera(test_bild, top_k=2)
|
||||
|
||||
print("\n🎯 Identifieringsresultat:")
|
||||
for k in kandidater:
|
||||
print(f" {k.lvx_id}: {k.namn_sv} (konfidens: {k.konfidens:.3f})")
|
||||
|
||||
# Spara modell
|
||||
model.spara(Path("/tmp/landvex-vision-model.pt"))
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,232 @@
|
||||
#!/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()
|
||||
Reference in New Issue
Block a user