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