/
nacha
/
pseado
Обзор
Документация
Войти
/
nacha
/
pseado
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
1
CI/CD
Аналитика
Безопасность
master
quality_checker/data_loader.py
62 строки
2 KB
alex
Add data loading and dataset class
07 авг 2026, 08:35
Верифицирован
07 авг 2026, 08:35
b437d3f
Код
Авторство
О чём код?
import os import cv2 import numpy as np import torch from torch.utils.data import Dataset from PIL import Image class DrawingDataset(Dataset): """ Загрузчик данных для обучения модели. Предполагается следующая структура: data/drawings/ images/ (jpg/png renderings) annotations/ (json с разметкой: bboxes, graph edges, labels) """ def __init__(self, root_dir: str, transform=None, mode="train"): self.root_dir = root_dir self.transform = transform self.images = sorted([f for f in os.listdir(os.path.join(root_dir, "images")) if f.endswith(('.jpg', '.png'))]) self.mode = mode def __len__(self): return len(self.images) def __getitem__(self, idx): img_name = self.images[idx] img_path = os.path.join(self.root_dir, "images", img_name) # Загрузка и ресайз img = Image.open(img_path).convert("RGB") if self.transform: img = self.transform(img) # Парсинг аннотаций (примерная логика) ann_path = os.path.join(self.root_dir, "annotations", img_name.replace('.jpg', '.json')) if os.path.exists(ann_path): # Формат: { "labels": [0, 2], "bbox": [x, y, w, h], "graph": [[0,1], [1,2]] } import json with open(ann_path) as f: ann = json.load(f) # Quality score (regression target) score = ann.get("score", 1.0) # Defect labels (multi-label) labels = torch.tensor(ann.get("labels", [0]*5), dtype=torch.float32) # Graph data edge_index = torch.tensor(ann.get("graph", []), dtype=torch.long).t().contiguous() x = torch.zeros(ann.get("num_nodes", 0), 10) # node features placeholder else: labels = torch.zeros(5, dtype=torch.float32) edge_index = torch.zeros((2, 0), dtype=torch.long) x = torch.zeros(0, 10) score = 0.0 return { "image": img, "labels": labels, "score": torch.tensor(score, dtype=torch.float32), "edge_index": edge_index, "node_features": x, "filename": img_name }