diff --git a/src/api/endpoints.py b/src/api/endpoints.py index 644300d..f346a33 100644 --- a/src/api/endpoints.py +++ b/src/api/endpoints.py @@ -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']}") diff --git a/src/classifiers/pet_breed_classifier.py b/src/classifiers/pet_breed_classifier.py index 0c6611e..807473c 100644 --- a/src/classifiers/pet_breed_classifier.py +++ b/src/classifiers/pet_breed_classifier.py @@ -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 diff --git a/src/core/orchestrator.py b/src/core/orchestrator.py index 336f033..cedc9ab 100644 --- a/src/core/orchestrator.py +++ b/src/core/orchestrator.py @@ -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. Сборка результата diff --git a/src/quality/medical_quality.py b/src/quality/medical_quality.py index f36dcb7..3a21b0b 100644 --- a/src/quality/medical_quality.py +++ b/src/quality/medical_quality.py @@ -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]: """ Комплексная проверка сегментации позвоночника diff --git a/src/segmentators/unet_segmentator.py b/src/segmentators/unet_segmentator.py index 04433e5..3030077 100644 --- a/src/segmentators/unet_segmentator.py +++ b/src/segmentators/unet_segmentator.py @@ -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)