api / app /services /segmenter.py
Omarelrayes's picture
Update app/services/segmenter.py
62c8927 verified
Raw
History Blame Contribute Delete
2.27 kB
# app/services/segmenter.py
import os
import io
import numpy as np
import tensorflow as tf
from PIL import Image
from typing import Dict, Any
os.environ["TF_ENABLE_ONEDNN_OPTS"] = "0"
from app.configs import get_segmentation_model, SEGMENTS_DIR
from app.core.preprocessing import preprocess_image
def segment_image(image_bytes: bytes, image_id: str = None) -> Dict[str, Any]:
"""
Run segmentation and save mask as PNG image.
Returns:
Dict with: status, masks_shape, max_confidence, mask_path, error
"""
model = get_segmentation_model()
if model is None:
return {
"status": "failed",
"error": "Segmentation model not loaded",
}
try:
# Preprocess
img_array = preprocess_image(image_bytes, target_size=(256, 256))
# Get original size for resizing mask
original_image = Image.open(io.BytesIO(image_bytes))
original_size = original_image.size
# Run model
result = model(tf.constant(img_array))
if isinstance(result, dict):
masks = list(result.values())[0].numpy()
else:
masks = result.numpy()
# Post-process mask
mask = masks[0] # Remove batch dimension
if mask.ndim == 3:
mask = mask[:, :, 0] # Take first channel
# Normalize to 0-255
mask = (mask * 255).astype(np.uint8)
# Create PIL image and resize to original size
mask_image = Image.fromarray(mask)
mask_image = mask_image.resize(original_size)
# 🔥 Save mask as PNG
mask_path = None
if image_id:
mask_filename = f"{image_id}_mask.png"
mask_path = str(SEGMENTS_DIR / mask_filename)
mask_image.save(mask_path, format='PNG')
print(f"💾 Mask saved to: {mask_path}")
return {
"status": "completed",
"masks_shape": list(masks.shape),
"max_confidence": float(np.max(masks)),
"mask_path": mask_path,
}
except Exception as e:
print(f"❌ SEGMENTATION FAILED: {e}")
return {
"status": "failed",
"error": str(e),
}