This commit is contained in:
denis 2026-09-02 21:06:28 +03:00
parent a35d796ac3
commit e891988644
5 changed files with 51 additions and 43 deletions

View File

@ -25,7 +25,7 @@ scorer = None
device = None
def load_models():
def load_models(path: str = "../../models/unet_cats_dogs.pth"):
"""Ленивая загрузка моделей"""
global model, scorer, device
@ -35,7 +35,9 @@ def load_models():
# Загрузка модели
model = UNet(in_channels=3, out_classes=2)
model_path = Path("../../models/unet_cats_dogs.pth")
cwd = Path.cwd()
print(f"Директория инициализации скрипта: #{cwd}")
model_path = Path(path)
if model_path.exists():
print(f"📦 Загрузка модели из {model_path}")
@ -61,7 +63,7 @@ async def analyze_image(file: UploadFile = File(...)):
"""Анализ изображения на качество"""
try:
# Загрузка модели
model, scorer, device = load_models()
# model, scorer, device = load_models()
# Чтение изображения
image_data = await file.read()
@ -75,47 +77,49 @@ async def analyze_image(file: UploadFile = File(...)):
original_image = np.array(image)
original_size = image.size # (width, height)
# Подготовка для модели
input_size = (256, 256)
image_resized = image.resize(input_size)
image_array = np.array(image_resized, dtype=np.float32) / 255.0
image_array = image_array.transpose(2, 0, 1)
image_tensor = torch.FloatTensor(image_array).unsqueeze(0).to(device)
# Инференс
with torch.no_grad():
output = model(image_tensor)
probs = torch.softmax(output, dim=1)
pred_mask = probs.argmax(dim=1).squeeze().cpu().numpy()
# Если маска пустая — пробуем порог
if pred_mask.sum() == 0:
print("⚠️ Маска пустая! Пробуем пороговую обработку...")
prob_class1 = probs[0, 1].cpu().numpy()
for threshold in [0.2, 0.15, 0.1]:
temp_mask = (prob_class1 > threshold).astype(np.int64)
if temp_mask.sum() > 0:
print(f" ✅ Найден объект при пороге {threshold}")
pred_mask = temp_mask
break
# Ресайз маски к оригинальному размеру
mask_pil = Image.fromarray(pred_mask.astype(np.uint8))
mask_pil = mask_pil.resize(original_size, resample=Image.NEAREST)
mask = np.array(mask_pil)
print(f" Финальная маска: {mask.shape}, сумма = {mask.sum()}")
result = orchestrator.analyze(image)
#
# # Подготовка для модели
# input_size = (256, 256)
# image_resized = image.resize(input_size)
# image_array = np.array(image_resized, dtype=np.float32) / 255.0
# image_array = image_array.transpose(2, 0, 1)
# image_tensor = torch.FloatTensor(image_array).unsqueeze(0).to(device)
#
# # Инференс
# with torch.no_grad():
# output = model(image_tensor)
# probs = torch.softmax(output, dim=1)
# pred_mask = probs.argmax(dim=1).squeeze().cpu().numpy()
#
# # Если маска пустая — пробуем порог
# if pred_mask.sum() == 0:
# print("⚠️ Маска пустая! Пробуем пороговую обработку...")
# prob_class1 = probs[0, 1].cpu().numpy()
# for threshold in [0.2, 0.15, 0.1]:
# temp_mask = (prob_class1 > threshold).astype(np.int64)
# if temp_mask.sum() > 0:
# print(f" ✅ Найден объект при пороге {threshold}")
# pred_mask = temp_mask
# break
#
# # Ресайз маски к оригинальному размеру
# mask_pil = Image.fromarray(pred_mask.astype(np.uint8))
# mask_pil = mask_pil.resize(original_size, resample=Image.NEAREST)
# mask = np.array(mask_pil)
#
# print(f" Финальная маска: {mask.shape}, сумма = {mask.sum()}")
# Конвертируем маску в base64 для передачи на фронтенд
mask_image = Image.fromarray((mask * 255).astype(np.uint8))
mask_image = Image.fromarray((result.get("mask") * 255).astype(np.uint8))
buffer = io.BytesIO()
mask_image.save(buffer, format='PNG')
mask_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
# Оценка качества
quality_scores = scorer.evaluate(image=original_image, segmentation=mask)
# quality_scores = scorer.evaluate(image=original_image, segmentation=result.get("mask"))
result = orchestrator.analyze(image)
quality_scores = result.get("quality")
print(f"\n📊 Результаты:")
print(f" Качество: {quality_scores['overall_quality']}")

View File

@ -50,8 +50,8 @@ class PetBreedClassifier(BaseClassifier):
# Заменяем последний слой на 37 классов (породы)
num_features = self.model.fc.in_features
self.model.fc = nn.Linear(num_features, 37)
model_path = "../../models/unet_cats_dogs.pth"
# self.model.fc = nn.Linear(num_features, 37)
model_path = "../models/unet_cats_dogs.pth"
# Загружаем обученные веса если есть
self.loaded = False
if model_path and Path(model_path).exists():
@ -135,9 +135,9 @@ class PetBreedClassifier(BaseClassifier):
return {
"breed_id": breed_id,
"breed_name": BREED_MAP.get(breed_id, "Unknown"),
# "breed_name": BREED_MAP.get(breed_id, "Unknown"),
"species_id": species_id,
"species_name": SPECIES_MAP.get(species_id, "Unknown"),
# "species_name": SPECIES_MAP.get(species_id, "Unknown"),
"confidence": confidence_score,
"loaded": True,
"error": None

View File

@ -76,8 +76,9 @@ class Orchestrator:
quality_result = self.quality_scorer.evaluate(
image=image_array,
segmentation=mask,
config=self.config
segmentation=mask
# ,
# config=self.config
)
# 4. Сборка результата

View File

@ -25,6 +25,9 @@ class MedicalQualityScorer(QualityScorer):
"medical_validation": medical_check
}
def check_femur_segmentation(self, segmentation: np.ndarray) -> Dict[str, Any]:
pass
def check_spine_segmentation(self, segmentation: np.ndarray) -> Dict[str, Any]:
"""
Комплексная проверка сегментации позвоночника

View File

@ -11,7 +11,7 @@ from src.segmentators.base import BaseSegmentator
class UNetSegmentator(BaseSegmentator):
"""Сегментатор на основе U-Net (для кошек/собак)"""
def __init__(self, model_path="models/unet_cats_dogs.pth", device='cpu'):
def __init__(self, model_path="../models/unet_cats_dogs.pth", device='cpu'):
self.device = device
self.model = UNet(in_channels=3, out_classes=2)