/
nacha
/
pseado
Обзор
Документация
Войти
/
nacha
/
pseado
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
1
CI/CD
Аналитика
Безопасность
master
quality_checker/model.py
125 строк
5 KB
alex
Add main neural network model architecture
07 авг 2026, 08:37
Верифицирован
07 авг 2026, 08:37
d0d6c1e
Код
Авторство
О чём код?
import torch import torch.nn as nn import torch.nn.functional as F from efficientnet_pytorch import EfficientNet class GNNLayer(nn.Module): """Simple Graph Neural Network Layer (GraphSAGE style)""" def __init__(self, in_channels, out_channels): super().__init__() self.linear = nn.Linear(in_channels, out_channels) self.edge_mlp = nn.Sequential( nn.Linear(in_channels * 2, in_channels), nn.ReLU(), nn.Linear(in_channels, out_channels) ) def forward(self, x, edge_index): # x: [num_nodes, in_channels] # edge_index: [2, num_edges] row, col = edge_index # Source node embeddings src = self.linear(x) dst = self.linear(x) # Message passing messages = self.edge_mlp(torch.cat([src[row], dst[col]], dim=1)) # Aggregation (Mean) out = x # Identity for residual connection logic if needed, here just update num_nodes = x.size(0) if num_nodes == 0: return torch.zeros((0, self.linear.out_features), device=x.device, dtype=x.dtype) row_sum = torch.zeros(num_nodes, dst.size(1), device=x.device) row_sum.scatter_add_(0, row.view(-1, 1).expand(-1, dst.size(1)), messages) deg_inv = 1.0 / torch.clamp(torch.bincount(row, minlength=num_nodes), min=1.0) deg_inv = deg_inv.view(-1, 1).expand(-1, dst.size(1)).to(x.device) return F.relu(row_sum * deg_inv) class DrawingQualityModel(nn.Module): """ Главная архитектура модели: - Backbone для визуальных черт - OCR энкодер для текстовых элементов - GNN для топологической проверки - Fusion Layer и Heads для предсказаний """ def __init__(self, cfg): super().__init__() self.cfg = cfg # 1. Визуальный энкодер (EfficientNet) self.backbone = EfficientNet.from_pretrained('efficientnet-b4') # Срезаем последний слой классификации self.backbone._fc = nn.Identity() visual_feat_dim = self.backbone._fc.in_features # 2. OCR Элемент (симуляция) # В реальности здесь DBNet + CRNN. Для стаба используем глобальные влипания self.text_feat_dim = 512 self.text_encoder = nn.Linear(self.text_feat_dim, self.text_feat_dim) # 3. GNN слой self.gnn = GNNLayer(self.text_feat_dim, self.text_feat_dim) # 4. Fusion Layer fusion_input = visual_feat_dim + self.text_feat_dim self.fusion = nn.Sequential( nn.Linear(fusion_input, 512), nn.ReLU(), nn.Dropout(0.5), nn.Linear(512, 256), nn.ReLU() ) # 5. Heads (Мультизадачность) # Head 1: Regress Quality Score (0..100) self.score_head = nn.Sequential( nn.Linear(256, 64), nn.ReLU(), nn.Linear(64, 1) ) # Head 2: Multi-label Defect Classification (5 classes) self.class_head = nn.Sequential( nn.Linear(256, 64), nn.ReLU(), nn.Linear(64, cfg.model.num_defect_classes) ) def forward(self, images, text_features, graph_data): """ Args: images: Tensor [B, 3, H, W] text_features: Tensor [B, T, C] (после энкодера OCR) graph_data: Dict with 'edge_index' and 'node_features' """ batch_size = images.size(0) # Визуальные признаки (Global Avg Pool) visual_feats = self.backbone(images) # Обработка текста и графа if text_features is not None and text_features.size(0) > 0: # Простейшая агрегация текста (Mean Pooling) text_agg = text_features.mean(dim=1) text_agg = self.text_encoder(text_agg) else: text_agg = torch.zeros(batch_size, self.text_feat_dim, device=images.device) # Fusion combined = torch.cat([visual_feats, text_agg], dim=1) fused = self.fusion(combined) # Предсказания score = torch.sigmoid(self.score_head(fused)) * 100.0 defects = torch.sigmoid(self.class_head(fused)) return { "score": score, "defects": defects }