develop
This commit is contained in:
parent
a35d796ac3
commit
e891988644
|
|
@ -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']}")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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. Сборка результата
|
||||
|
|
|
|||
|
|
@ -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]:
|
||||
"""
|
||||
Комплексная проверка сегментации позвоночника
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Reference in New Issue