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

View File

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

View File

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

View File

@ -25,6 +25,9 @@ class MedicalQualityScorer(QualityScorer):
"medical_validation": medical_check "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]: 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): class UNetSegmentator(BaseSegmentator):
"""Сегментатор на основе U-Net (для кошек/собак)""" """Сегментатор на основе 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.device = device
self.model = UNet(in_channels=3, out_classes=2) self.model = UNet(in_channels=3, out_classes=2)