Files
guilan-multimedia-lab/backend/processing/services.py

213 lines
7.4 KiB
Python

import time
import numpy as np
from django.conf import settings
from django.utils import timezone
from .algorithms import ProcessingError, average_images, decode_image, histogram, histogram_payload, normalize_to_uint8, process_image, verify_registration
from .models import ImageSession, ImageState, ProcessingJob
from .registry import apply_registered_operation
from .storage import load_image_array, payload_for_image, save_image_array
from .tasks import run_batch_job
def image_session_create(*, uploaded_file=None, image_base64=None):
if uploaded_file and uploaded_file.size > settings.MAX_UPLOAD_MB * 1024 * 1024:
raise ProcessingError(f"Upload exceeds {settings.MAX_UPLOAD_MB} MB.")
image = decode_image(uploaded_file=uploaded_file, base64_image=image_base64)
relative_path = save_image_array(image, "original")
hist = histogram(image)
session = ImageSession.objects.create(
original_image=relative_path,
width=image.shape[1],
height=image.shape[0],
channels=image.shape[2] if image.ndim == 3 else 1,
color_mode="RGB" if image.ndim == 3 else "L",
original_histogram=hist,
expires_at=timezone.now() + timezone.timedelta(hours=settings.IMAGE_SESSION_TTL_HOURS),
)
state = image_state_create(
session=session,
parent=None,
image=image,
operation="upload",
params={},
label="S0 Upload",
prefix="state-upload",
)
payload = {
"session_id": str(session.id),
"active_state_id": str(state.id),
"width": session.width,
"height": session.height,
"channels": session.channels,
"color_mode": session.color_mode,
"original_histogram": hist,
"histogram": state.histogram,
"expires_at": session.expires_at.isoformat(),
"states": [image_state_payload(state=state, include_image=True)],
}
payload.update(payload_for_image(image, relative_path))
return payload
def image_state_create(*, session, parent, image, operation, params, label=None, prefix="state"):
relative_path = save_image_array(image, prefix)
sequence = session.states.count()
state = ImageState.objects.create(
session=session,
parent=parent,
sequence=sequence,
label=label or f"S{sequence} {operation}",
operation=operation,
params=params or {},
image=relative_path,
width=image.shape[1],
height=image.shape[0],
channels=image.shape[2] if image.ndim == 3 else 1,
color_mode="RGB" if image.ndim == 3 else "L",
histogram=histogram_payload(image),
)
return state
def image_state_payload(*, state, include_image=True):
payload = {
"state_id": str(state.id),
"session_id": str(state.session_id),
"parent_state_id": str(state.parent_id) if state.parent_id else None,
"sequence": state.sequence,
"label": state.label,
"operation": state.operation,
"params": state.params,
"width": state.width,
"height": state.height,
"channels": state.channels,
"color_mode": state.color_mode,
"histogram": state.histogram,
"created_at": state.created_at.isoformat(),
}
if include_image:
image = load_image_array(state.image)
payload.update(payload_for_image(image, state.image))
else:
payload["image_path"] = state.image
payload["image_url"] = f"{settings.MEDIA_URL}{state.image}"
return payload
def image_states_payload(*, states):
return [image_state_payload(state=state, include_image=True) for state in states]
def image_state_apply_operation(*, state, operation, params):
if state.session.expired:
raise ProcessingError("Image session has expired.")
source = load_image_array(state.image)
result = apply_registered_operation(source, operation, params or {})
new_state = image_state_create(
session=state.session,
parent=state,
image=result,
operation=operation,
params=params or {},
label=None,
prefix=f"state-{operation}",
)
return image_state_payload(state=new_state, include_image=True)
def combine_states(*, states, operation, params=None):
params = params or {}
if len(states) < 2:
raise ProcessingError("At least two states are required.")
session = states[0].session
if any(state.session_id != session.id for state in states):
raise ProcessingError("All states must belong to the same session.")
images = [load_image_array(state.image) for state in states]
verify_registration(images)
if operation == "average":
result = average_images(images)
elif operation == "add":
result = np_clip_sum(images)
elif operation == "subtract":
result = normalize_to_uint8(images[0].astype("float32") - images[1].astype("float32"))
elif operation == "dot_product":
result = normalize_to_uint8(np.prod([image.astype("float32") / 255.0 for image in images], axis=0))
elif operation == "and":
result = images[0].copy()
for image in images[1:]:
result = result & image
elif operation == "or":
result = images[0].copy()
for image in images[1:]:
result = result | image
else:
raise ProcessingError(f"Unsupported combine operation '{operation}'.")
new_state = image_state_create(
session=session,
parent=states[0],
image=result,
operation=f"combine_{operation}",
params={**params, "state_ids": [str(state.id) for state in states]},
label=None,
prefix=f"state-combine-{operation}",
)
return image_state_payload(state=new_state, include_image=True)
def np_clip_sum(images):
total = np.zeros_like(images[0], dtype="float32")
for image in images:
total += image.astype("float32")
return np.clip(total, 0, 255).astype("uint8")
def image_session_process(*, session, operation, params):
if session.expired:
raise ProcessingError("Image session has expired.")
started_at = time.perf_counter()
source = load_image_array(session.original_image)
result = process_image(source, operation, params)
relative_path = save_image_array(result, f"processed-{operation}")
hist = histogram(result)
session.processed_image = relative_path
session.processed_histogram = hist
session.save(update_fields=["processed_image", "processed_histogram"])
payload = {
"session_id": str(session.id),
"operation": operation,
"params": params,
"processed_histogram": hist,
"elapsed_ms": round((time.perf_counter() - started_at) * 1000, 2),
}
payload.update(payload_for_image(result, relative_path))
return payload
def batch_job_create(*, operation, session_ids, params=None):
job = ProcessingJob.objects.create(operation=operation, params=params or {})
run_batch_job.delay(str(job.id), operation, [str(session_id) for session_id in session_ids])
return job
def processing_job_payload(*, job):
payload = {
"job_id": str(job.id),
"operation": job.operation,
"status": job.status,
"progress": job.progress,
"error": job.error,
"result_histogram": job.result_histogram,
}
if job.result_image:
image = load_image_array(job.result_image)
payload.update(payload_for_image(image, job.result_image))
return payload