382 lines
14 KiB
Python
382 lines
14 KiB
Python
import time
|
|
|
|
import cv2
|
|
import numpy as np
|
|
from django.conf import settings
|
|
from django.db.models import Max
|
|
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 delete_relative_file, load_fft_array, load_image_array, payload_for_image, save_fft_array, save_image_array
|
|
from .tasks import run_batch_job
|
|
|
|
|
|
def image_session_create(*, uploaded_file=None, image_base64=None):
|
|
"""Create a new image session and its initial S0 upload state.
|
|
|
|
The session groups all later processing states so a user can present the full workflow.
|
|
"""
|
|
|
|
if uploaded_file and uploaded_file.size > settings.MAX_UPLOAD_MB * 1024 * 1024:
|
|
raise ProcessingError(f"Upload exceeds {settings.MAX_UPLOAD_MB} MB.")
|
|
|
|
image = compact_workspace_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 compact_workspace_image(image):
|
|
"""Resize large uploads to the configured maximum dimension for faster processing.
|
|
|
|
This keeps classroom-sized experiments responsive even when the uploaded file is very large.
|
|
"""
|
|
|
|
max_dimension = int(getattr(settings, "IMAGE_WORKSPACE_MAX_DIMENSION", 1400))
|
|
if max_dimension <= 0:
|
|
return image
|
|
height, width = image.shape[:2]
|
|
longest = max(width, height)
|
|
if longest <= max_dimension:
|
|
return image
|
|
scale = max_dimension / float(longest)
|
|
next_size = (max(1, int(round(width * scale))), max(1, int(round(height * scale))))
|
|
return cv2.resize(image, next_size, interpolation=cv2.INTER_AREA).astype(np.uint8)
|
|
|
|
|
|
def image_state_create(*, session, parent, image, operation, params, label=None, prefix="state"):
|
|
"""Persist one processed image state with its metadata, parent, and histogram.
|
|
|
|
Saving every result makes it possible to compare steps and combine previous states later.
|
|
"""
|
|
|
|
relative_path = save_image_array(image, prefix)
|
|
max_sequence = session.states.aggregate(value=Max("sequence"))["value"]
|
|
sequence = 0 if max_sequence is None else max_sequence + 1
|
|
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):
|
|
"""Serialize an image state for the frontend workspace."""
|
|
|
|
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):
|
|
"""Serialize a list of image states."""
|
|
|
|
return [image_state_payload(state=state, include_image=True) for state in states]
|
|
|
|
|
|
def image_state_delete(*, state):
|
|
"""Delete a non-S0 state while keeping child states available.
|
|
|
|
Use it to remove unhelpful experiments without losing later useful results.
|
|
"""
|
|
|
|
if state.sequence == 0 or state.operation == "upload":
|
|
raise ProcessingError("The original S0 upload state cannot be deleted.")
|
|
image_path = state.image
|
|
fft_data_path = state.params.get("fft_data_path") if isinstance(state.params, dict) else None
|
|
state.children.update(parent=None)
|
|
state.delete()
|
|
delete_relative_file(image_path)
|
|
delete_relative_file(fft_data_path)
|
|
|
|
|
|
def image_state_apply_operation(*, state, operation, params):
|
|
"""Apply one registered algorithm to a selected state and save the result as a new state.
|
|
|
|
This is the main workspace action: every filter or transform becomes a reproducible step.
|
|
"""
|
|
|
|
if state.session.expired:
|
|
raise ProcessingError("Image session has expired.")
|
|
|
|
if operation == "fft_spectrum":
|
|
return image_state_fft_spectrum_create(state=state, params=params or {})
|
|
if operation == "inverse_fft_reconstruction":
|
|
return image_state_inverse_fft_create(state=state, params=params or {})
|
|
|
|
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 image_state_fft_spectrum_create(*, state, params):
|
|
"""Create an FFT visualization state and persist the actual complex spectrum.
|
|
|
|
The displayed image is only a spectrum preview; the saved FFT data is what the inverse step uses.
|
|
"""
|
|
|
|
source = load_image_array(state.image)
|
|
mode = params.get("mode", "log_magnitude")
|
|
if mode not in {"magnitude", "log_magnitude", "phase"}:
|
|
raise ProcessingError("FFT mode must be magnitude, log_magnitude, or phase.")
|
|
spectrum = image_fft(source)
|
|
preview = fft_preview_image(spectrum, mode)
|
|
fft_data_path = save_fft_array(spectrum, "state-fft-data")
|
|
operation_params = {
|
|
**params,
|
|
"mode": mode,
|
|
"fft_data_path": fft_data_path,
|
|
"source_state_id": str(state.id),
|
|
"source_color_mode": "RGB" if source.ndim == 3 else "L",
|
|
}
|
|
new_state = image_state_create(
|
|
session=state.session,
|
|
parent=state,
|
|
image=preview,
|
|
operation="fft_spectrum",
|
|
params=operation_params,
|
|
label=None,
|
|
prefix="state-fft_spectrum",
|
|
)
|
|
return image_state_payload(state=new_state, include_image=True)
|
|
|
|
|
|
def image_state_inverse_fft_create(*, state, params):
|
|
"""Create an image state by applying inverse FFT to a previously saved FFT state."""
|
|
|
|
if state.operation != "fft_spectrum" or not isinstance(state.params, dict) or not state.params.get("fft_data_path"):
|
|
raise ProcessingError("Inverse FFT Reconstruction must be applied to an FFT/DFT Spectrum View state.")
|
|
spectrum = load_fft_array(state.params["fft_data_path"])
|
|
result = inverse_fft_image(spectrum)
|
|
new_state = image_state_create(
|
|
session=state.session,
|
|
parent=state,
|
|
image=result,
|
|
operation="inverse_fft_reconstruction",
|
|
params=params or {},
|
|
label=None,
|
|
prefix="state-inverse_fft_reconstruction",
|
|
)
|
|
return image_state_payload(state=new_state, include_image=True)
|
|
|
|
|
|
def image_fft(image):
|
|
"""Return centered DFT data for grayscale or per-channel RGB images."""
|
|
|
|
if image.ndim == 2:
|
|
return np.fft.fftshift(np.fft.fft2(image.astype(np.float32)))
|
|
channels = [np.fft.fftshift(np.fft.fft2(image[:, :, idx].astype(np.float32))) for idx in range(image.shape[2])]
|
|
return np.stack(channels, axis=2)
|
|
|
|
|
|
def fft_preview_image(spectrum, mode):
|
|
"""Convert complex FFT data to a display-only uint8 preview image."""
|
|
|
|
if mode == "phase":
|
|
preview = np.angle(spectrum)
|
|
else:
|
|
preview = np.abs(spectrum)
|
|
if mode == "log_magnitude":
|
|
preview = np.log1p(preview)
|
|
return normalize_to_uint8(preview)
|
|
|
|
|
|
def inverse_fft_image(spectrum):
|
|
"""Apply ifft2(ifftshift(F)) to stored complex FFT data."""
|
|
|
|
def reconstruct_channel(channel):
|
|
reconstructed = np.real(np.fft.ifft2(np.fft.ifftshift(channel)))
|
|
return np.round(np.clip(reconstructed, 0, 255)).astype(np.uint8)
|
|
|
|
if spectrum.ndim == 2:
|
|
return reconstruct_channel(spectrum)
|
|
channels = [reconstruct_channel(spectrum[:, :, idx]) for idx in range(spectrum.shape[2])]
|
|
return np.stack(channels, axis=2)
|
|
|
|
|
|
def combine_states(*, states, operation, params=None):
|
|
"""Combine registered states using add, subtract, dot product, average, and/or.
|
|
|
|
Use it for MATLAB-like image arithmetic, change detection, masking, and K-image denoising.
|
|
"""
|
|
|
|
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):
|
|
"""Add several registered images and clip the result to [0, 255].
|
|
|
|
Use it to combine brightness/detail contributions while keeping the output displayable.
|
|
"""
|
|
|
|
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):
|
|
"""Run a legacy single-image operation against the original session image.
|
|
|
|
This keeps the older API working while the state workspace handles the main presentation flow.
|
|
"""
|
|
|
|
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):
|
|
"""Create a Celery-backed processing job for heavier batch operations.
|
|
|
|
Use it when an operation may take longer than an interactive request should block.
|
|
"""
|
|
|
|
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):
|
|
"""Serialize a processing job, including result image data when available."""
|
|
|
|
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
|