feat(reconstruction): Add asynchronous series reconstruction task with progress tracking

This commit is contained in:
Ross
2026-05-17 11:32:48 +01:00
parent c8d05818a4
commit 7ed0a8f378
4 changed files with 421 additions and 198 deletions
+248 -2
View File
@@ -3,11 +3,15 @@ from django.core.mail import send_mail
from django.http import HttpResponse
from django.shortcuts import get_object_or_404
from celery import shared_task
from atlas.models import Case
from atlas.models import Case, Series, SeriesImage
from generic.models import CimarCase
from rad.settings import REMOTE_URL, CIMAR_USERNAME, CIMAR_PASSWORD
from helpers.cimar import CimarAPI, NotFoundError
from pydicom.uid import generate_uid
from django.contrib.auth.models import User
from django.core.files.base import ContentFile
import copy
import io
@shared_task()
def push_case_to_cimar_task(case_id):
@@ -49,4 +53,246 @@ def push_case_to_cimar_task(case_id):
cimar_case.refresh_study()
return 10
return 10
@shared_task(bind=True)
def series_reconstruct_task(
self,
series_id,
user_id,
recon_planes,
slice_thickness_val=None,
slice_spacing_val=None,
recon_thickness_mode="mean",
):
"""Generate reconstructions asynchronously for a series with progress updates."""
import numpy as np
from loguru import logger
from atlas import views as atlas_views
series = get_object_or_404(Series, pk=series_id)
user = get_object_or_404(User, pk=user_id)
if not series.check_user_can_edit(user):
raise PermissionError("Permission denied")
images = list(series.get_images())
dicom_items = []
for image in images:
ds = atlas_views._read_series_image_dataset(image)
if ds is None:
continue
try:
arr = ds.pixel_array
if arr.ndim != 2:
continue
dicom_items.append((image, ds, arr))
except Exception:
continue
if len(dicom_items) < 2:
raise ValueError("Need at least 2 valid DICOM images in series for reconstruction")
base_shape = dicom_items[0][2].shape
dicom_items = [item for item in dicom_items if item[2].shape == base_shape]
if len(dicom_items) < 2:
raise ValueError("Not enough consistently-sized slices for reconstruction")
geom = atlas_views._extract_recon_geometry(dicom_items)
sorted_items = geom["sorted_items"]
volume = np.stack([item[2] for item in sorted_items], axis=0)
template_ds = sorted_items[0][1]
source_positions_mm = geom["source_positions_mm"]
origin_ipp = geom["origin_ipp"]
row_dir = geom["row_dir"]
col_dir = geom["col_dir"]
normal_dir = geom["normal_dir"]
native_row_spacing = geom["row_spacing"]
native_col_spacing = geom["col_spacing"]
native_z_spacing = geom["native_z_spacing"]
target_spacing = float(slice_spacing_val) if slice_spacing_val is not None else native_z_spacing
slab_thickness = float(slice_thickness_val) if slice_thickness_val is not None else target_spacing
volume_for_recon, recon_centers_mm = atlas_views._aggregate_volume_along_z(
volume,
source_positions_mm,
float(target_spacing),
float(slab_thickness),
recon_thickness_mode,
)
if volume_for_recon.shape[0] < 1:
raise ValueError("No reconstruction slices were generated")
base_z_offset = float(recon_centers_mm[0]) if len(recon_centers_mm) > 0 else 0.0
def build_position(origin, row_vector, col_vector, stack_vector, row_index=0, col_index=0, stack_offset=0.0):
pos = (
np.asarray(origin, dtype=float)
+ (np.asarray(stack_vector, dtype=float) * float(stack_offset))
+ (np.asarray(row_vector, dtype=float) * (float(row_index) * float(native_row_spacing)))
+ (np.asarray(col_vector, dtype=float) * (float(col_index) * float(native_col_spacing)))
)
return [float(pos[0]), float(pos[1]), float(pos[2])]
plane_slices = {}
for plane in recon_planes:
plane_norm = plane.lower()
if plane_norm == "axial":
recon_slices = [volume_for_recon[i, :, :] for i in range(volume_for_recon.shape[0])]
pixel_spacing_out = [native_row_spacing, native_col_spacing]
spacing_between_slices_out = float(target_spacing)
image_orientation_out = [
float(row_dir[0]),
float(row_dir[1]),
float(row_dir[2]),
float(col_dir[0]),
float(col_dir[1]),
float(col_dir[2]),
]
def position_for_index(idx):
return build_position(
origin_ipp,
row_dir,
col_dir,
normal_dir,
row_index=0,
col_index=0,
stack_offset=float(recon_centers_mm[idx]),
)
elif plane_norm == "coronal":
recon_slices = [volume_for_recon[:, i, :] for i in range(volume_for_recon.shape[1])]
pixel_spacing_out = [target_spacing, native_col_spacing]
spacing_between_slices_out = float(native_row_spacing)
image_orientation_out = [
float(normal_dir[0]),
float(normal_dir[1]),
float(normal_dir[2]),
float(col_dir[0]),
float(col_dir[1]),
float(col_dir[2]),
]
def position_for_index(idx):
return build_position(
origin_ipp,
row_dir,
col_dir,
normal_dir,
row_index=idx,
col_index=0,
stack_offset=base_z_offset,
)
elif plane_norm == "sagittal":
recon_slices = [volume_for_recon[:, :, i] for i in range(volume_for_recon.shape[2])]
pixel_spacing_out = [target_spacing, native_row_spacing]
spacing_between_slices_out = float(native_col_spacing)
image_orientation_out = [
float(normal_dir[0]),
float(normal_dir[1]),
float(normal_dir[2]),
float(row_dir[0]),
float(row_dir[1]),
float(row_dir[2]),
]
def position_for_index(idx):
return build_position(
origin_ipp,
row_dir,
col_dir,
normal_dir,
row_index=0,
col_index=idx,
stack_offset=base_z_offset,
)
else:
continue
plane_slices[plane_norm] = {
"recon_slices": recon_slices,
"pixel_spacing_out": pixel_spacing_out,
"spacing_between_slices_out": spacing_between_slices_out,
"image_orientation_out": image_orientation_out,
"position_for_index": position_for_index,
}
if not plane_slices:
raise ValueError("No valid reconstruction planes selected")
total_slices = sum(len(v["recon_slices"]) for v in plane_slices.values())
processed = 0
created_series = []
for plane_norm, cfg in plane_slices.items():
recon_series = atlas_views._create_series_derivative(series, f"Recon {plane_norm.title()}")
recon_series.series_instance_uid = generate_uid()
recon_series.save(update_fields=["series_instance_uid"])
for idx, arr2d in enumerate(cfg["recon_slices"]):
ds_new = copy.deepcopy(template_ds)
arr2d = np.asarray(arr2d, dtype=dicom_items[0][2].dtype)
ds_new.Rows = int(arr2d.shape[0])
ds_new.Columns = int(arr2d.shape[1])
ds_new.InstanceNumber = idx + 1
ds_new.SOPInstanceUID = generate_uid()
ds_new.SeriesInstanceUID = recon_series.series_instance_uid
ds_new.PixelData = arr2d.tobytes()
ds_new.PixelSpacing = [float(cfg["pixel_spacing_out"][0]), float(cfg["pixel_spacing_out"][1])]
ds_new.ImageOrientationPatient = cfg["image_orientation_out"]
ds_new.ImagePositionPatient = cfg["position_for_index"](idx)
ds_new.SliceThickness = float(slab_thickness)
ds_new.SpacingBetweenSlices = float(cfg["spacing_between_slices_out"])
out_io = io.BytesIO()
ds_new.save_as(out_io, write_like_original=False)
out_io.seek(0)
recon_image = SeriesImage(
series=recon_series,
position=idx + 1,
upload_filename=f"recon_{plane_norm}_{idx + 1}.dcm",
)
recon_image.image.save(
f"recon_{plane_norm}_{recon_series.pk}_{idx + 1}.dcm",
ContentFile(out_io.getvalue()),
save=False,
)
recon_image.save()
processed += 1
self.update_state(
state="PROGRESS",
meta={
"current": processed,
"total": total_slices,
"message": f"Generating {plane_norm} reconstruction ({processed}/{total_slices})",
},
)
created_series.append(
{
"id": recon_series.pk,
"url": recon_series.get_absolute_url(),
"description": recon_series.description or str(recon_series.pk),
}
)
logger.info(
"Reconstruction task complete for series {} with {} outputs",
series.pk,
len(created_series),
)
return {
"series_id": series.pk,
"created_series": created_series,
"target_spacing": float(target_spacing),
"slab_thickness": float(slab_thickness),
"mode": recon_thickness_mode,
}