NeuroSAM3 / scripts_build_samples.py
mmrech's picture
NeuroSAM3 v2.0: unified DICOM/NIfTI loader, real SAM 3 point/box prompts, fixed mask threshold, agent + pipelines rewrite, MCP tools
7214c3c verified
Raw History Blame Contribute Delete
5.16 kB
"""
Rebuild the bundled sample dataset from its public sources.
python scripts_build_samples.py [--out samples]
Sources (all openly distributed):
* Orthanc public demo server (orthanc.uclouvain.be/demo): BRAINIX FLAIR + T1-Gd MRI, PHENIX head CT (OsiriX teaching cases)
* Hugging Face dataset BJyotibrat/Masoud-Nickparvar-Brain-Tumor-MRI-Dataset (CC BY 4.0)
* Hugging Face dataset iraqigold/brain-stroke-ct-dataset (Teknofest 2021 brain stroke CT, with expert masks)
"""
import argparse, concurrent.futures as cf, glob, json, os, shutil, urllib.request, zipfile
import numpy as np
from PIL import Image
ORTHANC = "https://orthanc.uclouvain.be/demo"
SERIES = {"brainix_flair": ("1e2c125c-411b8e86-3f4fe68e-a7584dd3-c6da78f0", None),
"brainix_t1gd": ("dc0216d2-a406a5ad-31ef7a78-113ae9d9-29939f9e", slice(48, 96, 3)),
"phenix_headct": ("17cc7e52-4f1a3e4d-9182f727-56e9cc71-c037892f", slice(215, 322, 7))}
TUMOR = "https://huggingface.co/datasets/BJyotibrat/Masoud-Nickparvar-Brain-Tumor-MRI-Dataset"
STROKE = "https://huggingface.co/datasets/iraqigold/brain-stroke-ct-dataset/resolve/main/"
def get(url):
return json.load(urllib.request.urlopen(url, timeout=60))
def dl(url, path):
for _ in range(3):
try:
urllib.request.urlretrieve(url, path)
return path
except Exception as e:
err = e
raise err
def save_img(src, dst, maxside=512):
im = Image.open(src)
if im.mode == "RGBA":
im = Image.alpha_composite(Image.new("RGBA", im.size, (0, 0, 0, 255)), im)
im = im.convert("L")
if max(im.size) > maxside:
im.thumbnail((maxside, maxside))
im.save(dst, optimize=True)
def main(out):
tmp = os.path.join(out, "_tmp"); os.makedirs(tmp, exist_ok=True)
for d in ["mri_tumor", "mri_normal", "ct_hemorrhage/masks", "ct_ischemia/masks", "ct_normal", "nifti", "zips"] + [f"dicom/{k}" for k in SERIES]:
os.makedirs(os.path.join(out, d), exist_ok=True)
# --- DICOM series -------------------------------------------------
with cf.ThreadPoolExecutor(8) as ex:
for name, (sid, sl) in SERIES.items():
inst = [s[0] for s in get(f"{ORTHANC}/series/{sid}/ordered-slices")["SlicesShort"]]
if sl:
inst = inst[sl]
list(ex.map(lambda a: dl(f"{ORTHANC}/instances/{a[1]}/file", os.path.join(out, "dicom", name, f"{name}_{a[0]:03d}.dcm")), enumerate(inst, 1)))
with zipfile.ZipFile(os.path.join(out, "zips", f"{name}_series.zip"), "w", zipfile.ZIP_DEFLATED) as z:
for f in sorted(glob.glob(os.path.join(out, "dicom", name, "*.dcm"))):
z.write(f, arcname=os.path.basename(f))
# --- NIfTI from FLAIR --------------------------------------------
import pydicom, nibabel as nib
fl = sorted(glob.glob(os.path.join(out, "dicom/brainix_flair/*.dcm")))
dss = [pydicom.dcmread(f) for f in fl]
vol = np.stack([d.pixel_array.astype(np.int16) for d in dss], axis=-1)
ps = [float(x) for x in dss[0].PixelSpacing]
st = abs(float(dss[1].ImagePositionPatient[2]) - float(dss[0].ImagePositionPatient[2]))
arr = np.transpose(vol, (1, 0, 2))[::-1, ::-1, :]
nib.save(nib.Nifti1Image(arr, np.diag([ps[1], ps[0], st, 1.0])), os.path.join(out, "nifti/brainix_flair.nii.gz"))
# --- Tumor MRI (CC BY 4.0) ----------------------------------------
for cls, outn in [("glioma", "glioma"), ("meningioma", "meningioma"), ("pituitary", "pituitary"), ("notumor", "normal")]:
files = [x["path"] for x in get(f"{TUMOR.replace('huggingface.co/datasets', 'huggingface.co/api/datasets')}/tree/main/Testing/{cls}")][5:8]
for i, p in enumerate(files, 1):
src = dl(f"{TUMOR}/resolve/main/{p}", os.path.join(tmp, os.path.basename(p)))
save_img(src, os.path.join(out, "mri_normal" if cls == "notumor" else "mri_tumor", f"{outn}_{i:02d}.png"))
# --- Stroke CT with masks -----------------------------------------
fs = [s["rfilename"] for s in get("https://huggingface.co/api/datasets/iraqigold/brain-stroke-ct-dataset")["siblings"]]
picks = {"ct_hemorrhage": ("hemorrhage", [f for f in fs if f.startswith("Bleeding/images")][:5], True),
"ct_ischemia": ("ischemia", [f for f in fs if f.startswith("Ischemia/images")][:2], True),
"ct_normal": ("normal", [f for f in fs if f.startswith("Normal/images")][40:43], False)}
for folder, (stem, files, has_mask) in picks.items():
for i, p in enumerate(files, 1):
save_img(dl(STROKE + p, os.path.join(tmp, "img.png")), os.path.join(out, folder, f"{stem}_{i:02d}.png"))
if has_mask:
m = Image.open(dl(STROKE + p.replace("images", "masks"), os.path.join(tmp, "mask.png"))).convert("L").point(lambda v: 255 if v > 127 else 0)
m.save(os.path.join(out, folder, "masks", f"{stem}_{i:02d}_mask.png"), optimize=True)
shutil.rmtree(tmp, ignore_errors=True)
print("samples built in", out)
if __name__ == "__main__":
ap = argparse.ArgumentParser(); ap.add_argument("--out", default="samples"); a = ap.parse_args(); main(a.out)