diff --git a/.github/python-versions.json b/.github/python-versions.json index b9254d0..3f62ac7 100644 --- a/.github/python-versions.json +++ b/.github/python-versions.json @@ -1,4 +1,9 @@ { - "test": ["3.11", "3.12"], + "test": [ + "3.11", + "3.12", + "3.13", + "3.14" + ], "docs": "3.12" } diff --git a/.github/workflows/dependency_watch.yml b/.github/workflows/dependency_watch.yml index 4e70627..dd2fc77 100644 --- a/.github/workflows/dependency_watch.yml +++ b/.github/workflows/dependency_watch.yml @@ -49,6 +49,7 @@ jobs: env: GH_TOKEN: ${{ github.token }} run: | + set -o pipefail BRANCH="${{ steps.watch.outputs.branch }}" # A pull request in any state for this branch is a standing answer. if [ "$(gh pr list --repo "$GITHUB_REPOSITORY" --head "$BRANCH" \ diff --git a/CITATION.cff b/CITATION.cff new file mode 100644 index 0000000..2917c8c --- /dev/null +++ b/CITATION.cff @@ -0,0 +1,11 @@ +cff-version: 1.2.0 +message: "If you use this software, please cite it as below." +title: "MBIRTorch: High-performance tomographic reconstruction using PyTorch" +authors: + - family-names: Buzzard + given-names: Gregery T. + - family-names: Bouman + given-names: Charles A. +year: 2026 +url: "https://github.com/cabouman/mbirtorch" +license: BSD-3-Clause diff --git a/ci/dependency_watch.py b/ci/dependency_watch.py index e828396..8ed7f35 100644 --- a/ci/dependency_watch.py +++ b/ci/dependency_watch.py @@ -21,6 +21,11 @@ import urllib.request CPU_INDEX_URL = "https://download.pytorch.org/whl/cpu/torch/" +# The Python versions GitHub's hosted runners can install (setup-python's +# source of truth). A version torch supports but the runners lack cannot +# be tested and must not be proposed. +RUNNER_MANIFEST_URL = ("https://raw.githubusercontent.com/actions/" + "python-versions/main/versions-manifest.json") REMOTE_RAW = "https://raw.githubusercontent.com/cabouman/mbirtorch/prerelease/" VERSION_FILE = ".github/python-versions.json" PYPROJECT = "pyproject.toml" @@ -76,6 +81,22 @@ def parse_torch_index(html): return newest, sorted(files[newest], key=lambda v: int(v.split(".")[1])) +def parse_runner_manifest(text): + """The Python versions GitHub's runners install, as minors like "3.12". + Only entries marked stable count; release candidates do not make a + version testable.""" + minors = set() + for entry in json.loads(text): + if not entry.get("stable"): + continue + parts = str(entry.get("version", "")).split(".") + if len(parts) >= 2 and parts[0] == "3" and parts[1].isdigit(): + minors.add(f"3.{parts[1]}") + if not minors: + raise ValueError("no stable Python versions found in the runner manifest") + return minors + + def parse_version_file(text): """The matrix from the version file. Returns (test_list, docs_version).""" data = json.loads(text) @@ -101,12 +122,18 @@ def _minor(v): return tuple(int(x) for x in v.split(".")[:2]) -def divergence(torch_release, torch_list, matrix, python_floor, torch_floor): - """The divergence, as a dict. Versions below the Python floor are - reported informationally and never proposed.""" +def divergence(torch_release, torch_list, matrix, python_floor, torch_floor, + runner_minors=None): + """The divergence, as a dict. Versions below the Python floor, and + versions absent from ``runner_minors`` (the versions GitHub's runners + install), are reported informationally and never proposed.""" below_floor = [v for v in torch_list if _minor(v) < _minor(python_floor)] eligible = [v for v in torch_list if _minor(v) >= _minor(python_floor)] additions = [v for v in eligible if v not in matrix] + not_on_runners = [] + if runner_minors is not None: + not_on_runners = [v for v in additions if v not in runner_minors] + additions = [v for v in additions if v in runner_minors] removals = [v for v in matrix if v not in torch_list] torch_newest_minor = ".".join(str(x) for x in _minor(torch_release)) torch_advance = (torch_newest_minor @@ -118,6 +145,7 @@ def divergence(torch_release, torch_list, matrix, python_floor, torch_floor): "python_floor": python_floor, "torch_floor": torch_floor, "below_floor": below_floor, + "not_on_runners": not_on_runners, "additions": additions, "removals": removals, "torch_advance": torch_advance, @@ -255,6 +283,15 @@ def main(argv=None): torch_release, torch_list = parse_torch_index(fetch(CPU_INDEX_URL)) print(f"dependency-watch: torch {torch_release} supports {torch_list}") + try: + runner_minors = parse_runner_manifest(fetch(RUNNER_MANIFEST_URL)) + except (OSError, urllib.error.URLError, ValueError) as e: + print(f"dependency-watch: RUNNER MANIFEST NOT READ " + f"({RUNNER_MANIFEST_URL}): {e}") + print("dependency-watch: verdict UNKNOWN (cannot tell which versions " + "the runners install; this is not 'no divergence')") + return 1 + try: matrix, docs_version = parse_version_file(read(vf_source)) except (OSError, urllib.error.URLError) as e: @@ -265,13 +302,17 @@ def main(argv=None): print(f"dependency-watch: matrix {matrix}, docs {docs_version} ({vf_source})") python_floor, torch_floor = parse_pyproject(read(pp_source)) - d = divergence(torch_release, torch_list, matrix, python_floor, torch_floor) + d = divergence(torch_release, torch_list, matrix, python_floor, torch_floor, + runner_minors=runner_minors) if args.json: print(json.dumps(d, indent=2)) if d["below_floor"]: print(f"dependency-watch: below the {python_floor} floor, not proposed: " f"{d['below_floor']}") + if d["not_on_runners"]: + print(f"dependency-watch: torch supports but GitHub runners do not " + f"install yet, not proposed: {d['not_on_runners']}") if d["any"]: print(f"dependency-watch: DIVERGENCE -> branch {branch_name(d)}") if d["additions"]: diff --git a/ci/test_dependency_watch.py b/ci/test_dependency_watch.py index 73d061d..60ce07e 100644 --- a/ci/test_dependency_watch.py +++ b/ci/test_dependency_watch.py @@ -9,7 +9,8 @@ sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) from dependency_watch import (parse_torch_index, parse_version_file, - parse_pyproject, divergence, branch_name) + parse_pyproject, parse_runner_manifest, + divergence, branch_name) INDEX_HTML = """ @@ -57,6 +58,30 @@ def test_parse_pyproject_floors(): assert torch_floor == "2.13" +RUNNER_MANIFEST_JSON = """[ + {"version": "3.15.0-rc.2", "stable": false}, + {"version": "3.14.2", "stable": true}, + {"version": "3.14.0", "stable": true}, + {"version": "3.13.9", "stable": true}, + {"version": "3.12.12", "stable": true}, + {"version": "3.11.14", "stable": true} +]""" + + +def test_parse_runner_manifest_stable_minors_only(): + minors = parse_runner_manifest(RUNNER_MANIFEST_JSON) + assert minors == {"3.11", "3.12", "3.13", "3.14"} # 3.15 rc excluded + + +def test_version_on_torch_index_but_not_on_runners_is_not_proposed(): + d = divergence("2.13.0", ["3.11", "3.12", "3.13", "3.14", "3.15"], + ["3.11", "3.12"], "3.11", "2.13", + runner_minors={"3.11", "3.12", "3.13", "3.14"}) + assert d["not_on_runners"] == ["3.15"] # informational only + assert d["additions"] == ["3.13", "3.14"] + assert branch_name(d) == "nightly/python-matrix-add-3.13-3.14" + + def test_multi_version_addition_with_below_floor_exclusion(): d = divergence("2.13.0", ["3.10", "3.11", "3.12", "3.13", "3.14"], ["3.11", "3.12"], "3.11", "2.13") diff --git a/demo/demo_1_parallel_basics.py b/demo/demo_1_parallel_basics.py new file mode 100644 index 0000000..b28ede6 --- /dev/null +++ b/demo/demo_1_parallel_basics.py @@ -0,0 +1,42 @@ +"""Demo 1: the basic MBIRTorch pipeline. + +Make a simple 3D phantom, forward project it to get a sinogram, and +reconstruct it with model-based iterative reconstruction (MBIR). + +In a real application you would skip the phantom and load your measured +sinogram as a numpy array with axes in the order +(views, detector rows, detector channels). +""" + +import numpy as np +import mbirtorch + +# Problem size: small enough to run on a laptop CPU in about a minute. +num_views = 128 +num_det_rows = 128 +num_det_channels = 128 + +# Make a phantom and project it to get a synthetic sinogram. +phantom, sinogram, params = mbirtorch.generate_demo_data( + model_type='parallel', object_type='shepp-logan', + num_views=num_views, num_det_rows=num_det_rows, + num_det_channels=num_det_channels) + +# The generator also returns the projection angles it used. +angles = params['angles'] + +# Build the reconstruction model from the sinogram shape and the angles. +ct_model = mbirtorch.ParallelBeamModel(sinogram.shape, angles) + +# Reconstruct. Everything is at its default value. The one parameter worth +# trying first is sharpness (default 1.0): higher gives crisper edges, lower +# gives smoother images. To change it: ct_model.set_params(sharpness=1.5) +recon, recon_dict = ct_model.recon(sinogram) + +# Compare the reconstruction to the phantom. +nrmse = np.linalg.norm(recon - phantom) / np.linalg.norm(phantom) +print(f'Normalized RMS error between reconstruction and phantom: {nrmse:.3f}') + +# View them side by side. Use the sliders to change slice and intensity. +mbirtorch.slice_viewer(phantom, recon, data_dicts=[None, recon_dict], vmin=0.0, + title='Phantom (left) and MBIR reconstruction (right)') diff --git a/demo/demo_1_shepp_logan.py b/demo/demo_1_shepp_logan.py deleted file mode 100644 index b54d25a..0000000 --- a/demo/demo_1_shepp_logan.py +++ /dev/null @@ -1,84 +0,0 @@ -"""Demo 1: 3D Shepp-Logan reconstruction with mbirtorch (the mbirjax demo_1 -equivalent). - -Generates a Shepp-Logan phantom, forward projects it to a sinogram, and runs -the VCD reconstruction; prints the per-iteration traces and the final NRMSE -against the phantom. Run parameters sit at the top (no CLI arguments); -MODEL_TYPE selects the geometry ('parallel' or 'cone', as in the mbirjax -demo). Set SHOW_SLICES = True to explore the ground truth phantom and the -reconstruction in the slice viewer (the recon's data dict rides along). -""" - -import time - -import numpy as np - -import mbirtorch - -# ── run parameters ──────────────────────────────────────────────────────────── -MODEL_TYPE = "cone" # 'parallel' or 'cone' -SINOGRAM_SHAPE = (80, 100, 128) # (num_views, num_det_rows, num_det_channels) -MAX_ITERATIONS = 15 -SHARPNESS = 1.0 -# Devices are not named here. The model resolves cuda > mps > cpu on its -# own, and on a machine with several CUDA devices it spreads the -# reconstruction across the ones that can hold their share. To pin it, call -# model.configure_devices(num_devices=1) or configure_devices(devices=['cpu']). -SEED = 0 -SHOW_SLICES = True -# ────────────────────────────────────────────────────────────────────────────── - - -def build_model(): - n_views, _, num_channels = SINOGRAM_SHAPE - if MODEL_TYPE == "cone": - # Cone beam: full-circle angles, and source-detector / source-iso - # distances in the goldens' convention (magnification 2). The auto - # recon geometry sets the recon shape, including the axial padding. - angles = np.linspace(0, 2 * np.pi, n_views, endpoint=False) - _model = mbirtorch.ConeBeamModel( - SINOGRAM_SHAPE, angles, - source_detector_dist=4 * num_channels, - source_iso_dist=2 * num_channels) - return _model - if MODEL_TYPE == "parallel": - angles = np.linspace(0, np.pi, n_views, endpoint=False) - return mbirtorch.ParallelBeamModel(SINOGRAM_SHAPE, angles) - raise ValueError(f"MODEL_TYPE must be 'parallel' or 'cone', got {MODEL_TYPE!r}") - - -def main(): - model = build_model() - model.set_params(no_warning=True, sharpness=SHARPNESS) - recon_shape = model.get_params("recon_shape") - print(f"model = {MODEL_TYPE}, device = {model.torch_device}, " - f"recon_shape = {recon_shape}") - - phantom = mbirtorch.generate_3d_shepp_logan_low_dynamic_range(recon_shape) - sinogram = model.forward_project(phantom) - weights = mbirtorch.gen_weights(sinogram / np.max(sinogram), - weight_type="transmission_root") - - np.random.seed(SEED) - t0 = time.time() - recon, recon_dict = model.recon(sinogram, weights=weights, - max_iterations=MAX_ITERATIONS) - elapsed = time.time() - t0 - - nrmse = float(np.linalg.norm(recon - phantom) / np.linalg.norm(phantom)) - rp = recon_dict["recon_params"] - print(f"\nElapsed: {elapsed:.2f} s for {rp['num_iterations']} iterations") - print(f"Final forward loss: {rp['fm_rmse'][-1]:.4f}") - print(f"NRMSE vs phantom: {nrmse:.4f}") - mbirtorch.get_memory_stats() - - if SHOW_SLICES: - mbirtorch.slice_viewer( - phantom, recon, - slice_label=["ground truth phantom", "mbirtorch recon"], - data_dicts=[None, recon_dict], - title=f"Shepp-Logan {MODEL_TYPE} demo (NRMSE {nrmse:.4f})") - - -if __name__ == "__main__": - main() diff --git a/demo/demo_2_cone_beam.py b/demo/demo_2_cone_beam.py new file mode 100644 index 0000000..610f297 --- /dev/null +++ b/demo/demo_2_cone_beam.py @@ -0,0 +1,76 @@ +"""Demo 2: cone-beam reconstruction, with the practices real data needs. + +This demo adds four things to the basic pipeline of demo 1: + +1. Cone-beam geometry, which needs two distances: source to detector, and + source to the rotation axis. +2. Simulated measurement noise with the physically correct structure: + rays through dense material are noisier. +3. Noise weighting: the weights tell the reconstruction to trust the + noisier measurements less. +4. Saving the reconstruction to a file. +""" + +import numpy as np +import mbirtorch + +# Problem size. +num_views = 128 +num_det_rows = 128 +num_det_channels = 128 + +# Make a phantom and its cone-beam sinogram. target_max_attenuation scales +# the phantom so the sinogram is in attenuation units (the units of real +# -log(I/I0) data), roughly in the range [0, 6]. +phantom, sinogram, params = mbirtorch.generate_demo_data( + model_type='cone', object_type='shepp-logan', + num_views=num_views, num_det_rows=num_det_rows, + num_det_channels=num_det_channels, target_max_attenuation=6.0) + +# Add measurement noise. For a transmission scan with a dosage of +# lambda_0 input photons per measurement, the attenuation measurements are +# approximately +# y = ybar + sqrt(exp(ybar) / lambda_0) * W, W ~ N(0, 1), +# so the noise standard deviation grows with attenuation. (Bouman and +# Sauer, "A Unified Approach to Statistical Tomography Using Coordinate +# Descent Optimization," IEEE Trans. on Image Processing, 1996.) +dosage = 10000.0 +noise_std = np.sqrt(np.exp(sinogram) / dosage) +rng = np.random.default_rng(0) +sinogram = sinogram + noise_std * rng.standard_normal(sinogram.shape).astype(np.float32) + +# The generator also returns the geometry it used. +angles = params['angles'] +source_detector_dist = params['source_detector_dist'] +source_iso_dist = params['source_iso_dist'] + +# Build the cone-beam model. The two distances set the cone geometry. +ct_model = mbirtorch.ConeBeamModel(sinogram.shape, angles, + source_detector_dist=source_detector_dist, + source_iso_dist=source_iso_dist) + +# Noise weights. The noise model above has variance exp(y) / lambda_0, so +# down-weighting by the transmission gives the noisier measurements less +# influence. For a first look at any new data set, weights=None is also fine. +weights = mbirtorch.gen_weights(sinogram, weight_type='transmission_root') + +# Sharpness is the main image-quality control: higher gives crisper edges, +# lower gives smoother images. Typical useful range is about -1 to 2. +ct_model.set_params(sharpness=1.0) + +# Reconstruct. +recon, recon_dict = ct_model.recon(sinogram, weights=weights) + +nrmse = np.linalg.norm(recon - phantom) / np.linalg.norm(phantom) +print(f'Normalized RMS error between reconstruction and phantom: {nrmse:.3f}') + +# View the phantom and the reconstruction side by side. +mbirtorch.slice_viewer(phantom, recon, data_dicts=[None, recon_dict], vmin=0.0, + title='Phantom (left) and cone-beam MBIR reconstruction (right)') + +# Save the reconstruction and its settings to one file. The file can be +# reloaded later for viewing, or to continue from this result: +# recon, recon_dict = mbirtorch.TomographyModel.load_recon_hdf5(filepath) +filepath = './output/demo2_recon.h5' +ct_model.save_recon_hdf5(filepath, recon, recon_dict) +print(f'Reconstruction saved to {filepath}') diff --git a/demo/demo_3_parallel_roi.py b/demo/demo_3_parallel_roi.py new file mode 100644 index 0000000..6b72988 --- /dev/null +++ b/demo/demo_3_parallel_roi.py @@ -0,0 +1,68 @@ +"""Demo 3: a region-of-interest scan (object extends outside the field of view). + +In many real parallel-beam applications the object is wider than the +detector, so only a region of interest is scanned. Voxels outside the +field of view still contribute to some measurements, and if the +reconstruction ignores them, their contributions get pushed into the image +as artifacts. + +The fix is to enlarge the reconstruction region a little (about 1.3 times +the field of view), giving those outside contributions somewhere to go. +This demo reconstructs without and with the enlargement so you can see the +artifacts and their fix. +""" + +import numpy as np +import mbirtorch + +# Problem size. +num_views = 128 +num_det_rows = 128 +num_det_channels = 128 + +sinogram_shape = (num_views, num_det_rows, num_det_channels) +angles = np.linspace(0, np.pi, num_views, endpoint=False) + +# Make a phantom 1.5 times wider than the field of view in both lateral +# directions, and project it. The generation model is told the phantom's +# true size; only its projection onto the detector is kept. +gen_model = mbirtorch.ParallelBeamModel(sinogram_shape, angles) +phantom_shape = (int(1.5 * num_det_channels), int(1.5 * num_det_channels), + num_det_rows) +phantom = mbirtorch.generate_3d_shepp_logan_low_dynamic_range(phantom_shape) +gen_model.set_params(recon_shape=phantom_shape) +sinogram = gen_model.forward_project(phantom) + +# Reconstruction 1: the default region (exactly the field of view). +# The outside contributions have nowhere to go, so artifacts appear. +model_default = mbirtorch.ParallelBeamModel(sinogram_shape, angles) +recon_default, dict_default = model_default.recon(sinogram) + +# Reconstruction 2: enlarge the region by 1.3 in both lateral directions. +model_padded = mbirtorch.ParallelBeamModel(sinogram_shape, angles) +model_padded.scale_recon_shape(row_scale=1.3, col_scale=1.3) +recon_padded, dict_padded = model_padded.recon(sinogram) + +# Compare both to the phantom over the same central region of interest. +def center_crop(volume, rows, cols): + r0 = (volume.shape[0] - rows) // 2 + c0 = (volume.shape[1] - cols) // 2 + return volume[r0:r0 + rows, c0:c0 + cols, :] + +rows, cols = recon_default.shape[0], recon_default.shape[1] +phantom_roi = center_crop(phantom, rows, cols) +padded_roi = center_crop(recon_padded, rows, cols) + +nrmse_default = (np.linalg.norm(recon_default - phantom_roi) + / np.linalg.norm(phantom_roi)) +nrmse_padded = (np.linalg.norm(padded_roi - phantom_roi) + / np.linalg.norm(phantom_roi)) +print(f'Region-of-interest error without enlargement: {nrmse_default:.3f}') +print(f'Region-of-interest error with enlargement: {nrmse_padded:.3f}') + +# View: phantom region, the artifacted reconstruction, and the fixed one. +mbirtorch.slice_viewer( + phantom_roi, recon_default, padded_roi, + data_dicts=[None, dict_default, dict_padded], vmin=0.0, vmax=1.0, + title='Phantom region (left), default recon with artifacts (center),\n' + 'enlarged-region recon (right)') diff --git a/demo/demo_4_cone_axial_fov.py b/demo/demo_4_cone_axial_fov.py new file mode 100644 index 0000000..2d1d3f8 --- /dev/null +++ b/demo/demo_4_cone_axial_fov.py @@ -0,0 +1,82 @@ +"""Demo 4: cone-beam artifacts when the object extends past the top and +bottom of the field of view. + +In a cone-beam scan the X-rays diverge, so material just above and below +the field of view is still measured in many views. If the reconstruction +stops exactly at the field of view, those measurements corrupt the top and +bottom slices. + +The fix is the axial_pad_fraction parameter, which extends the +reconstruction axially so the outside material has somewhere to go. This +demo reconstructs without and with the padding. +""" + +import numpy as np +import mbirtorch + +# Problem size. +num_views = 128 +num_det_rows = 128 +num_det_channels = 128 + +sinogram_shape = (num_views, num_det_rows, num_det_channels) +angles = np.linspace(0, 2 * np.pi, num_views, endpoint=False) +source_detector_dist = 4 * num_det_channels +source_iso_dist = 2 * num_det_channels + +# Make a phantom 1.5 times taller than the field of view, and project it. +gen_model = mbirtorch.ConeBeamModel(sinogram_shape, angles, + source_detector_dist=source_detector_dist, + source_iso_dist=source_iso_dist) +fov_shape = tuple(gen_model.get_params('recon_shape')) +phantom_shape = (fov_shape[0], fov_shape[1], int(1.5 * fov_shape[2])) +phantom = mbirtorch.generate_3d_shepp_logan_low_dynamic_range(phantom_shape) +gen_model.set_params(recon_shape=phantom_shape) +sinogram = gen_model.forward_project(phantom) + +# Reconstruction 1: the default region. The slices near the top and +# bottom are corrupted by the material outside the field of view. +model_default = mbirtorch.ConeBeamModel(sinogram_shape, angles, + source_detector_dist=source_detector_dist, + source_iso_dist=source_iso_dist) +recon_default, dict_default = model_default.recon(sinogram) + +# Reconstruction 2: extend the region axially. axial_pad_fraction=1.0 +# pads each end far enough to cover every measured ray. After changing a +# geometry parameter, call auto_set_recon_geometry() to recompute the +# reconstruction region. +model_padded = mbirtorch.ConeBeamModel(sinogram_shape, angles, + source_detector_dist=source_detector_dist, + source_iso_dist=source_iso_dist) +model_padded.set_params(axial_pad_fraction=1.0) +model_padded.auto_set_recon_geometry() +recon_padded, dict_padded = model_padded.recon(sinogram) + +# Compare both to the phantom over the slices of the field of view. +def center_slices(volume, num_slices): + s0 = (volume.shape[2] - num_slices) // 2 + return volume[:, :, s0:s0 + num_slices] + +num_fov_slices = recon_default.shape[2] +phantom_fov = center_slices(phantom, num_fov_slices) +padded_fov = center_slices(recon_padded, num_fov_slices) + +# The artifacts concentrate in the slices near the ends, so measure there: +# the top and bottom eighth of the field of view. +n_end = max(1, num_fov_slices // 8) +ends = list(range(n_end)) + list(range(num_fov_slices - n_end, num_fov_slices)) + +def end_error(recon): + diff = recon[:, :, ends] - phantom_fov[:, :, ends] + return np.linalg.norm(diff) / np.linalg.norm(phantom_fov[:, :, ends]) + +print(f'End-slice error without axial padding: {end_error(recon_default):.3f}') +print(f'End-slice error with axial padding: {end_error(padded_fov):.3f}') + +# View all three. Look at the top and bottom slices, where the difference +# is largest. +mbirtorch.slice_viewer( + phantom_fov, recon_default, padded_fov, + data_dicts=[None, dict_default, dict_padded], vmin=0.0, slice_axis=1, + title='Phantom (left), default recon with axial artifacts (center),\n' + 'axially padded recon (right)') diff --git a/demo/demo_5_direct_vs_mbir.py b/demo/demo_5_direct_vs_mbir.py new file mode 100644 index 0000000..582ce3d --- /dev/null +++ b/demo/demo_5_direct_vs_mbir.py @@ -0,0 +1,54 @@ +"""Demo 5: direct reconstruction (FBP) versus model-based reconstruction (MBIR). + +Filtered back projection (FBP) is fast and works well when there are many +views and little noise. Model-based iterative reconstruction (MBIR) costs +more computation but stays good when the data gets hard. + +This demo runs both methods twice: once with plenty of views, where FBP is +serviceable, and once with very few views, where FBP breaks down and the +MBIR advantage is unmistakable. The sparse case is meant to be +illustrative, not practical. +""" + +import numpy as np +import mbirtorch + +# Problem size. +num_det_rows = 128 +num_det_channels = 128 + +def make_data(num_views): + phantom, sinogram, params = mbirtorch.generate_demo_data( + model_type='parallel', object_type='shepp-logan', + num_views=num_views, num_det_rows=num_det_rows, + num_det_channels=num_det_channels) + return phantom, sinogram, params['angles'] + +def nrmse(recon, phantom): + return np.linalg.norm(recon - phantom) / np.linalg.norm(phantom) + +# Case 1: plenty of views (128). FBP is serviceable; MBIR is cleaner. +phantom, sinogram, angles = make_data(num_views=128) +model = mbirtorch.ParallelBeamModel(sinogram.shape, angles) +fbp_many = model.direct_recon(sinogram) +mbir_many, _ = model.recon(sinogram) +print(f'128 views: FBP error {nrmse(fbp_many, phantom):.3f}, ' + f'MBIR error {nrmse(mbir_many, phantom):.3f}') + +# Case 2: very few views (16). FBP produces streaks; MBIR holds up. +phantom, sinogram, angles = make_data(num_views=16) +model = mbirtorch.ParallelBeamModel(sinogram.shape, angles) +fbp_sparse = model.direct_recon(sinogram) +mbir_sparse, mbir_sparse_dict = model.recon(sinogram) +print(f'16 views: FBP error {nrmse(fbp_sparse, phantom):.3f}, ' + f'MBIR error {nrmse(mbir_sparse, phantom):.3f}') + +# View the sparse-view case: phantom, FBP, MBIR. +mbirtorch.slice_viewer( + phantom, fbp_sparse, mbir_sparse, + data_dicts=[None, None, mbir_sparse_dict], vmin=0.0, + title='16 views: phantom (left), FBP with streaks (center), MBIR (right)') + +# The practical rule: with many views and low noise, FBP is fast and good +# enough, and MBIR uses it internally as a starting point. With few views, +# high noise, or metal, MBIR is worth its computation. diff --git a/demo/demo_6_helical.py b/demo/demo_6_helical.py new file mode 100644 index 0000000..8db1e32 --- /dev/null +++ b/demo/demo_6_helical.py @@ -0,0 +1,52 @@ +"""Demo 6: helical cone-beam reconstruction. + +In a helical scan the object moves steadily along the rotation axis while +the source rotates, so a short detector can cover a long object. Each view +therefore has an axial shift as well as an angle. The pitch is the travel +per rotation divided by the detector height, so a pitch of 1.0 means the +object moves one detector height per rotation. +""" + +import numpy as np +import mbirtorch + +# Problem size. The detector is short (few rows); the helix covers a +# longer object. +num_views = 180 +num_det_rows = 32 +num_det_channels = 128 + +helical_pitch = 1.0 # travel per rotation, in detector heights +helical_z_range = 40.0 # total travel over the scan, in ALU + +# Make a phantom and its helical sinogram. +phantom, sinogram, params = mbirtorch.generate_demo_data( + model_type='cone', object_type='shepp-logan', + num_views=num_views, num_det_rows=num_det_rows, + num_det_channels=num_det_channels, + use_helical=True, helical_pitch=helical_pitch, + helical_z_range=helical_z_range) + +# The per-view axial shifts are what make the scan helical. +angles = params['angles'] +helical_z_shifts = params['helical_z_shifts'] +source_detector_dist = params['source_detector_dist'] +source_iso_dist = params['source_iso_dist'] + +# Build the model. A helical scan uses the ordinary cone-beam model with +# one addition: the per-view axial shifts. +ct_model = mbirtorch.ConeBeamModel(sinogram.shape, angles, + source_detector_dist=source_detector_dist, + source_iso_dist=source_iso_dist, + helical_z_shifts=helical_z_shifts) + +# Reconstruct. Note that the reconstruction covers the full helical +# travel, so it has many more slices than the detector has rows. +recon, recon_dict = ct_model.recon(sinogram) +print(f'Detector rows: {num_det_rows}; reconstruction slices: {recon.shape[2]}') + +nrmse = np.linalg.norm(recon - phantom) / np.linalg.norm(phantom) +print(f'Normalized RMS error between reconstruction and phantom: {nrmse:.3f}') + +mbirtorch.slice_viewer(phantom, recon, data_dicts=[None, recon_dict], vmin=0.0, + title='Phantom (left) and helical reconstruction (right)') diff --git a/demo/demo_7_multiaxis.py b/demo/demo_7_multiaxis.py new file mode 100644 index 0000000..47ca0bd --- /dev/null +++ b/demo/demo_7_multiaxis.py @@ -0,0 +1,40 @@ +"""Demo 7: the multiaxis parallel geometry (laminography). + +In this geometry each view has two angles: the usual rotation about the +vertical axis (the azimuth), plus a tilt of the beam out of the horizontal +plane (the elevation). A constant tilt is laminography, which is useful +for flat objects such as circuit boards. +""" + +import numpy as np +import mbirtorch + +# Problem size and tilt. +num_views = 120 +num_det_rows = 96 +num_det_channels = 128 +elevation_degrees = 30.0 + +# Make a phantom and its tilted sinogram. +phantom, sinogram, params = mbirtorch.generate_demo_data( + model_type='multiaxis', elevation_degrees=elevation_degrees, + num_views=num_views, num_det_rows=num_det_rows, + num_det_channels=num_det_channels) + +# View the sinogram. Each view looks at the object from 30 degrees above +# the horizontal, so the projections show the object at an angle. +mbirtorch.slice_viewer(sinogram, slice_axis=0, slice_label='View', vmin=0.0, + title=f'Sinogram at {elevation_degrees:.0f} degree tilt') + +# The generator also returns the (azimuth, elevation) angle pairs it used. +angles = params['angles'] + +# Build the model and reconstruct. +ct_model = mbirtorch.MultiAxisParallelModel(sinogram.shape, angles) +recon, recon_dict = ct_model.recon(sinogram) + +nrmse = np.linalg.norm(recon - phantom) / np.linalg.norm(phantom) +print(f'Normalized RMS error between reconstruction and phantom: {nrmse:.3f}') + +mbirtorch.slice_viewer(phantom, recon, data_dicts=[None, recon_dict], vmin=0.0, + title='Phantom (left) and laminography reconstruction (right)') diff --git a/demo/demo_8_units_and_voxels.py b/demo/demo_8_units_and_voxels.py new file mode 100644 index 0000000..aeaf6a1 --- /dev/null +++ b/demo/demo_8_units_and_voxels.py @@ -0,0 +1,67 @@ +"""Demo 8: units, detector spacing, and voxel shape. + +MBIRTorch measures every length in ALUs (arbitrary length units): you pick +one physical unit — say millimeters — and use it for every distance and +spacing. By default the detector spacing is 1 ALU, and the reconstruction +uses cubic voxels sized to match the detector. + +The one rule this demo drives home: after changing any geometry parameter, +call auto_set_recon_geometry() so the reconstruction geometry is recomputed. +Setting the parameter alone is not enough. +""" + +import numpy as np +import mbirtorch + +# A small cone-beam model. All lengths below are in ALUs; if your detector +# spacing is 0.2 mm and you work in mm, you would enter 0.2. +num_views = 128 +num_det_rows = 128 +num_det_channels = 128 +sinogram_shape = (num_views, num_det_rows, num_det_channels) +angles = np.linspace(0, 2 * np.pi, num_views, endpoint=False) + +ct_model = mbirtorch.ConeBeamModel(sinogram_shape, angles, + source_detector_dist=4 * num_det_channels, + source_iso_dist=2 * num_det_channels) + +print('Default geometry (detector spacing 1 ALU):') +print(f" recon_shape = {tuple(ct_model.get_params('recon_shape'))}, " + f"delta_voxel = {ct_model.get_params('delta_voxel'):.3f} ALU") + +# Change the detector spacing. THE RULE: this alone does not update the +# reconstruction geometry -- the voxel size below is now stale. +ct_model.set_params(delta_det_channel=0.5, delta_det_row=0.5) +print('After set_params(delta_det_channel=0.5) alone (STALE):') +print(f" delta_voxel = {ct_model.get_params('delta_voxel'):.3f} ALU " + '<- unchanged, wrong') + +# Recompute the geometry. Now the voxel size follows the new spacing. +ct_model.auto_set_recon_geometry() +print('After auto_set_recon_geometry():') +print(f" recon_shape = {tuple(ct_model.get_params('recon_shape'))}, " + f"delta_voxel = {ct_model.get_params('delta_voxel'):.3f} ALU") + +# Restore the default spacing. The rule applies to every change. +ct_model.set_params(delta_det_channel=1.0, delta_det_row=1.0) +ct_model.auto_set_recon_geometry() + +# Voxel shape. Voxels need not be cubes: voxel_slice_aspect = 2.0 makes +# each slice twice as thick as the in-plane voxel size, halving the number +# of slices (useful when axial resolution matters less than memory). +ct_model.set_params(voxel_slice_aspect=2.0) +ct_model.auto_set_recon_geometry() # the same rule again +recon_shape = tuple(ct_model.get_params('recon_shape')) +print(f'With voxel_slice_aspect = 2.0: recon_shape = {recon_shape}') + +# Reconstruct with the thick slices and view the result. +phantom = mbirtorch.generate_3d_shepp_logan_low_dynamic_range(recon_shape) +sinogram = ct_model.forward_project(phantom) +recon, recon_dict = ct_model.recon(sinogram) + +nrmse = np.linalg.norm(recon - phantom) / np.linalg.norm(phantom) +print(f'Thick-slice reconstruction error vs its phantom: {nrmse:.3f}') + +mbirtorch.slice_viewer(phantom, recon, data_dicts=[None, recon_dict], vmin=0.0, + title='Thick-slice (voxel_slice_aspect = 2) phantom ' + 'and reconstruction') diff --git a/demo/demo_9_denoiser.py b/demo/demo_9_denoiser.py new file mode 100644 index 0000000..014a985 --- /dev/null +++ b/demo/demo_9_denoiser.py @@ -0,0 +1,30 @@ +"""Demo 9: the qGGMRF denoiser. + +The denoiser is the reconstruction's image model used on its own: it takes +a noisy 3D image and returns a smoothed one that preserves edges. No +geometry or sinogram is involved. Its one knob is sigma_noise, your +estimate of the noise standard deviation: larger values smooth more. +""" + +import numpy as np +import mbirtorch + +# Make a clean phantom and add noise with a known standard deviation. +shape = (128, 128, 128) +noise_std = 0.1 +phantom = mbirtorch.generate_3d_shepp_logan_low_dynamic_range(shape) +noisy = phantom + noise_std * np.random.default_rng(0).standard_normal(shape).astype(np.float32) + +# Denoise. Try sigma_noise above and below the true noise level to see +# over- and under-smoothing. +denoiser = mbirtorch.QGGMRFDenoiser(shape) +denoised, denoise_dict = denoiser.denoise(noisy, sigma_noise=noise_std) + +def nrmse(image): + return np.linalg.norm(image - phantom) / np.linalg.norm(phantom) + +print(f'Error of the noisy image: {nrmse(noisy):.3f}') +print(f'Error of the denoised image: {nrmse(denoised):.3f}') + +mbirtorch.slice_viewer(noisy, denoised, data_dicts=[None, denoise_dict], vmin=0.0, + title='Noisy image (left) and qGGMRF denoised image (right)') diff --git a/dev_scripts/refresh_widening_floors.py b/dev_scripts/refresh_widening_floors.py index 9c22aa3..edea726 100644 --- a/dev_scripts/refresh_widening_floors.py +++ b/dev_scripts/refresh_widening_floors.py @@ -163,11 +163,20 @@ def build_plan(smoke=False): def unmeasured_families(): - """Floor families a model class declares that the table has no rows for. + """Model classes whose automatic device count is set by floors that were + never measured for them, keyed by the floor family each class declares. - A geometry added without a measurement silently inherits the parallel - floors; this is where that shows up as work to do rather than as a - surprise in someone's log. + A geometry reaches this state two ways, and both are work to do rather + than a surprise in someone's log: + + * it DECLARES a ``_floor_family`` the table has no rows for, keyed here + under that name; or + * it declares no family at all -- the inherited base value -- and so + falls back to the ``wf.DEFAULT_FAMILY`` floors, keyed here under None. + + The second case is the one a newly ported geometry arrives in, and it is + the one this function used to skip, which left the tool silent about + exactly the classes relying on the fallback. """ import mbirtorch from mbirtorch.tomography_model import TomographyModel @@ -175,12 +184,20 @@ def unmeasured_families(): seen, known = {}, set(wf.families()) for name in dir(mbirtorch): cls = getattr(mbirtorch, name) - family = (getattr(cls, '_floor_family', None) - if isinstance(cls, type) - and issubclass(cls, TomographyModel) else None) - if family is not None and family not in known: - seen.setdefault(family, []).append(name) - return seen + if not (isinstance(cls, type) and issubclass(cls, TomographyModel)): + continue + # The base class is not a geometry, and a subclass that does not + # reconstruct through the shared VCD loop never reaches the automatic + # device-count decision the floors govern -- QGGMRFDenoiser subclasses + # TomographyModel but refuses recon, so no floor ever applies to it. + if cls is TomographyModel or cls.recon is not TomographyModel.recon: + continue + family = getattr(cls, '_floor_family', None) + if family is None or family not in known: + # An exported alias and its class are the same object, so record + # the class's own name once rather than once per exported name. + seen.setdefault(family, set()).add(cls.__name__) + return {family: sorted(names) for family, names in seen.items()} def print_plan(plan, smoke): @@ -217,13 +234,21 @@ def print_plan(plan, smoke): 'cell.'.format(family, count)) missing = unmeasured_families() if missing: - for family, classes in sorted(missing.items()): - print(' NEEDS MEASUREMENT: floor family {!r} is declared by {} ' - 'but has no rows; it currently inherits the {} floors.' - .format(family, ', '.join(sorted(classes)), - wf.DEFAULT_FAMILY)) + # The None key sorts first: a class that declares no family is taking + # the fallback silently, which is the case worth reading first. + for family, classes in sorted( + missing.items(), + key=lambda item: (item[0] is not None, item[0] or '')): + if family is None: + print(' NEEDS MEASUREMENT: {} declare no floor family, so the ' + '{} floors govern their automatic device count.' + .format(', '.join(classes), wf.DEFAULT_FAMILY)) + else: + print(' NEEDS MEASUREMENT: floor family {!r} is declared by {} ' + 'but has no rows; it currently inherits the {} floors.' + .format(family, ', '.join(classes), wf.DEFAULT_FAMILY)) else: - print(' every declared floor family has rows.') + print(' every model class is governed by floors measured for it.') # ── the worker: one arm, one subprocess ────────────────────────────────────── @@ -238,9 +263,17 @@ def _build_model(family, cell, device): model = mbirtorch.ConeBeamModel( tuple(cell), angles, source_detector_dist=4.0 * num_channels, source_iso_dist=2.0 * num_channels) - else: + elif family == 'parallel': angles = np.linspace(0, np.pi, num_views, endpoint=False) model = mbirtorch.ParallelBeamModel(tuple(cell), angles) + else: + # Falling through to parallel beam here would time parallel beam and + # record the result under this family's name, which is the one way a + # floor can be wrong without anything looking wrong. + raise ValueError( + 'refresh_widening_floors cannot build a model for floor family ' + '{!r}. Add its geometry to _build_model before measuring it.' + .format(family)) if device != 'cuda': # CPU/MPS only: the env pin is a CUDA mechanism (the policy # short-circuits below two visible devices), so the smoke has to place diff --git a/docs/source/credits.rst b/docs/source/credits.rst index 25d4cba..cd45294 100644 --- a/docs/source/credits.rst +++ b/docs/source/credits.rst @@ -9,7 +9,7 @@ The MBIR Development Team is listed below in alphabetical order: **MBIRTorch Sponsors** -We would like to thank the following sponsors for their financial support in the development of both this python package: +We would like to thank the following sponsors for their financial support in the development of this python package: * Eli Lilly Company * Oak Ridge National Laboratory @@ -17,17 +17,22 @@ We would like to thank the following sponsors for their financial support in the **Citation** -MBIRTorch is a PyTorch port of `MBIRJAX `__. Please use the following Bibtex citation when referencing this software. :: - @Misc {mbirjax-2024, - author = {Charles A. Bouman, Gregery T. Buzzard, Mingqi Yang, Ziyun Li, Diyu Yang, M. Samin Chowdhury, Karl Weisenburger, Caden Cardell, Brendt Wohlberg, and Chen Zhang}, - title = {{MBIRJAX}: {H}igh-performance tomographic reconstruction}, - howpublished = {Software library available from \url{https://github.com/cabouman/mbirjax}}, - year = 2024 + @misc{mbirtorch-2026, + title = {{MBIRTorch}: {H}igh-performance tomographic reconstruction using {PyTorch}}, + author = {Gregery T. Buzzard and Charles A. Bouman}, + howpublished = {Software library available from \url{https://github.com/cabouman/mbirtorch}}, + year = 2026 } +Alternatively, GitHub's "Cite this repository" button on the repository page +generates this citation from the repository's ``CITATION.cff`` file. + +MBIRTorch is a PyTorch port of `MBIRJAX `__ +:cite:`mbirjax-2024`; please also cite it when referencing the underlying methods. + **References** .. bibliography:: diff --git a/docs/source/demos_and_faqs.rst b/docs/source/demos_and_faqs.rst index 163e354..a430e51 100644 --- a/docs/source/demos_and_faqs.rst +++ b/docs/source/demos_and_faqs.rst @@ -7,16 +7,34 @@ Demos and FAQs Demos ----- -The basic demo below illustrates some of the features of MBIRTorch: - -* **Basic Demo:** `Python script `__ - -Follow the installation instructions in :ref:`InstallationDocs` and run the script directly. - -Then adjust some of the parameters to better understand how the code works. -If you have a GPU, you can increase the problem size by changing ``num_views``, ``num_det_rows``, and ``num_det_channels``. - -There are more demos here: `MBIRTorch demos `__ +The demo scripts are in the `demo folder `__. +Follow the installation instructions in :ref:`InstallationDocs`, then run any script directly. +Each is short and self-contained; adjust the parameters near the top and rerun to see their effect. + +.. list-table:: + :header-rows: 1 + :widths: 34 66 + + * - Script + - What it demonstrates + * - ``demo_1_parallel_basics.py`` + - The basic pipeline: make a phantom, project it to a sinogram, reconstruct, view. + * - ``demo_2_cone_beam.py`` + - Cone-beam geometry, simulated measurement noise, noise weighting, saving results. + * - ``demo_3_parallel_roi.py`` + - Region-of-interest reconstruction when the object extends outside the field of view. + * - ``demo_4_cone_axial_fov.py`` + - Cone-beam artifacts from material above and below the field of view, and axial padding. + * - ``demo_5_direct_vs_mbir.py`` + - Direct reconstruction (FBP) versus model-based reconstruction (MBIR), including sparse views. + * - ``demo_6_helical.py`` + - Helical cone-beam scanning and reconstruction. + * - ``demo_7_multiaxis.py`` + - The multiaxis parallel geometry (laminography): tilted views and their reconstruction. + * - ``demo_8_units_and_voxels.py`` + - Physical units (ALUs), detector spacing, voxel shape, and auto_set_recon_geometry(). + * - ``demo_9_denoiser.py`` + - The qGGMRF denoiser applied to a noisy 3D image. Data Generation @@ -40,8 +58,9 @@ geometry, so you can try MBIRTorch without a real dataset: Key options: * ``object_type`` -- ``'shepp-logan'`` or ``'cube'``. -* ``model_type`` -- ``'parallel'`` or ``'cone'``; ``params`` returns the matching - geometry parameters (always the view ``angles``, plus the source distances for cone beam). +* ``model_type`` -- ``'parallel'``, ``'cone'``, or ``'multiaxis'``; ``params`` returns the + matching geometry parameters (always the view ``angles``, plus the source distances for + cone beam and the tilt for multiaxis). * ``num_views``, ``num_det_rows``, ``num_det_channels`` -- the sinogram size; increase these (with a GPU) to make a larger problem. * ``target_max_attenuation`` -- scales the phantom so its sinogram has a realistic peak attenuation @@ -69,8 +88,7 @@ You can improve the reconstruction by increasing recon_shape: Note that the scale factor need only be large enough to give some padding around the region of valid projection -- it does not need to match the size of the true object. Larger scale factors will lead to increased time and memory. -.. PENDING(demos): mbirjax adds "See Demo 2: Large Object for an example of this." after the - first paragraph above. Restore that sentence when a matching mbirtorch demo exists. +See ``demo_3_parallel_roi.py`` for an example of this. Q: Why is my reconstruction blurry? +++++++++++++++++++++++++++++++++++ @@ -106,10 +124,6 @@ also make sure axial padding is disabled (``axial_pad_fraction=0``, the default We continue to improve the time and memory efficiency of MBIRTorch. -.. PENDING(demos): mbirjax closes that paragraph with "In either case, you can do a center - cropped reconstruction as in Demo 3: Cropped Center, although as seen in that demo, this - can introduce an intensity shift and other artifacts." Restore when a matching demo exists. - Q: Why does my reconstruction have artifacts? +++++++++++++++++++++++++++++++++++++++++++++ @@ -123,9 +137,10 @@ close to the object. For transmission tomography, it is critically important to preprocess the raw photon measurements by normalizing by an air-scan and taking the negative log of the ratio. We provide simple preprocessing utilities in ``mbirtorch.preprocess`` for doing this, and we plan to provide more utilities for specific instruments in the future. -In conebeam scans, it is sometimes the case that the rotation direction is reversed. -This can cause the reconstruction to look blurry or distorted. -You can correct this by simply taking the negative of your view angles. +In cone-beam scans, it is sometimes the case that the rotation direction is reversed. +The symptom is a reconstruction that is subtly warped, with shapes distorted and the top and +bottom of the object mirrored. You can correct this by taking the negative of your view +angles, or equivalently reversing their order with ``angles[::-1]``. A common artifact is rings near the center of the reconstruction that are generated when the center-of-rotation is not in the center of the detector. This can be corrected by setting the parameter ``det_channel_offset`` to reposition @@ -159,10 +174,6 @@ bad detector pixels and ``remove_sino_offset`` for a residual sinogram offset. A bright ring at the outer *boundary* of the reconstruction -- typically accompanied by the "Lateral FoV truncation detected" warning -- means the object extends past the field of view; see the next FAQ. -.. PENDING(demos): mbirjax adds "See Demo 3: Wrong Rotation Direction above for an example of - what can happen if the rotation direction is incorrect." to the rotation-direction - paragraph. Restore when a matching mbirtorch demo exists. - Q: What does the "Lateral FoV truncation detected" warning mean? ++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ diff --git a/docs/source/dev_maintenance.rst b/docs/source/dev_maintenance.rst new file mode 100644 index 0000000..94f78d1 --- /dev/null +++ b/docs/source/dev_maintenance.rst @@ -0,0 +1,57 @@ +Package Maintenance +=================== + +The following describes procedures for basic package maintenance. + +Unit Tests +---------- + +From the repository root, in the ``mbirtorch`` conda environment:: + + python -m pytest -n 4 tests ci + +To include the cross-framework parity tests against mbirjax, first generate +the golden archives (``tests/generate_goldens.py``, run in the mbirjax +environment), then:: + + python -m pytest -m "goldens or not goldens" tests + +The same tests run automatically on every push and pull request. + +Releasing a New Version +----------------------- + +This is only available for registered maintainers. + +1. Update ``__version__`` in ``mbirtorch/__init__.py`` and merge to + ``prerelease``. This is the only place the version number is written. + +2. On GitHub, draft a new release: tag ``vX.Y.ZrcN``, target ``prerelease``, + check "Set as a pre-release", and publish. This uploads to TestPyPI. + +3. Check the TestPyPI upload:: + + dev_scripts/check_published_wheel.sh --testpypi --version X.Y.ZrcN + +4. Open a pull request from ``prerelease`` to ``main`` and merge it when the + checks pass. + +5. Draft a new release: tag ``vX.Y.Z``, target ``main``, and publish. Then + approve the ``pypi`` environment on the workflow run page. This uploads + to PyPI. + +6. Check the PyPI upload:: + + dev_scripts/check_published_wheel.sh --version X.Y.Z + +The documentation rebuilds automatically: ``latest`` follows ``main``, and +``stable`` follows the highest release tag. + +Notes +----- + +* Uploads use PyPI Trusted Publishing; no token or password is stored. +* The tested Python versions are in ``.github/python-versions.json``. A + nightly check opens a pull request when torch's supported versions change; + merging it is the whole update. +* Manual upload with ``twine`` remains available as a fallback. diff --git a/docs/source/dev_sharding_overview.rst b/docs/source/dev_sharding_overview.rst index 21e7675..a649f3b 100644 --- a/docs/source/dev_sharding_overview.rst +++ b/docs/source/dev_sharding_overview.rst @@ -129,16 +129,27 @@ and ``sum_band_to_owner`` (reduce-scatter). Broadcast-to-N is the transpose of sum-from-N, which is what keeps forward and back projection adjoint under sharding. +The reduce **streams**. It forms the running total for a band once on the +slice-owner and then adds each arriving partial one bounded row slab at a time, +so the owner holds one slab per source above that total instead of every +partial at once. The summation order is untouched, so the streamed result is +bit for bit the one-shot sum. This is what makes the reduce shrink as devices +are added: it used to hold n whole bands, and n bands of 1/n of the volume each +is the same number of bytes at every device count. + **The default band is the whole shard**, which differs from MBIRJAX deliberately and on measurement. MBIRJAX's sweeps found time flat across band length, so it -streams by default for the memory win. The torch banded pass is instead -orchestration-bound, because the fan-out per band is eager: a sub-band default -measured 47 to 66 percent more warm reconstruction time at the two-device cells. +streams by default for the memory win. The torch banded pass pays a fixed +orchestration cost per band: with the compiled kernels in place, sub-band walks +measured 2 to 23 percent more busy time at parallel 1024 with two devices, +depending on the walk (an earlier pre-kernel reading of 47 to 66 percent +overstated the cost). MBIRJAX's stream-even-at-one-device rationale is also void here, because a single torch device never runs the banded drivers at all -- the trivial path uses the plain projectors. A smaller band remains a real **memory** lever, since the -per-band broadcast copy, the per-band partial, and each slice-owner's reduce -gather all scale with it. Set ``forward_project_slice_band`` or +per-band broadcast copy and the per-band partial scale with it; what it sets in +the reduce is the running total the slabs are added into, the bands already +reduced this pass being held either way. Set ``forward_project_slice_band`` or ``back_project_slice_band`` on the model to opt in (``_slice_band_length`` in ``tomography_model.py``). diff --git a/docs/source/index.rst b/docs/source/index.rst index 1236866..df4df4d 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -99,12 +99,10 @@ MBIRTorch: High-performance tomographic reconstruction dev_sharding_overview dev_projector_kernels dev_api + dev_maintenance .. PENDING(dashboard): restore dev_performance_dashboard to the toctree above when that page lands (held at Greg's request). -.. PENDING(maintenance): restore dev_maintenance to the toctree above when the - release workflow is implemented and the page is rewritten around it - (release_workflow.md, "Releasing, once set up"). .. _PyTorch: https://pytorch.org/docs/stable/index.html diff --git a/docs/source/refs.bib b/docs/source/refs.bib index e9090f4..3f12f19 100644 --- a/docs/source/refs.bib +++ b/docs/source/refs.bib @@ -32,3 +32,17 @@ @inproceedings{2024CV4SciencePoster } + +@misc{mbirjax-2024, + title={{MBIRJAX}: {H}igh-performance tomographic reconstruction}, + author={Bouman, Charles A and Buzzard, Gregery T and Yang, Mingqi and Li, Ziyun and Yang, Diyu and Chowdhury, M Samin and Weisenburger, Karl and Cardell, Caden and Wohlberg, Brendt and Zhang, Chen}, + howpublished={Software library available from \url{https://github.com/cabouman/mbirjax}}, + year={2024} +} + +@misc{mbirtorch-2026, + title={{MBIRTorch}: {H}igh-performance tomographic reconstruction using {PyTorch}}, + author={Buzzard, Gregery T and Bouman, Charles A}, + howpublished={Software library available from \url{https://github.com/cabouman/mbirtorch}}, + year={2026} +} diff --git a/docs/source/usr_api.rst b/docs/source/usr_api.rst index 665f4a3..ff3c7d7 100644 --- a/docs/source/usr_api.rst +++ b/docs/source/usr_api.rst @@ -35,7 +35,7 @@ individual pages for more detail. See :ref:`DemosFAQs` for examples. _sharding.run_per_device, get_psf_radii). Restoring these options therefore requires a different mechanism, not a narrower - __all__. The measured numbers are in plans/torch_port/docs.md. + __all__. .. automodule:: mbirtorch :no-index: diff --git a/docs/source/usr_multi_gpu.rst b/docs/source/usr_multi_gpu.rst index f091f98..5e45a1e 100644 --- a/docs/source/usr_multi_gpu.rst +++ b/docs/source/usr_multi_gpu.rst @@ -101,8 +101,9 @@ Tips for efficiency gather back to the host, so it can feed another on-device step directly. * **Trade memory for time with a smaller band.** Setting ``forward_project_slice_band`` or ``back_project_slice_band`` on the model streams the slice axis in smaller pieces. This is - a memory lever: on a measured 4-device 512-cell run it took peak memory from 6.6 GiB to - 2.6 GiB for about 8 percent more time. Leave it unset unless a run is memory-constrained. + a memory lever: a measured 2-device run at the 1024 class saved about 0.5 GB of per-device + peak for about 2 percent more time at the 252-slice band, and narrower bands saved slightly + more memory for more time. Leave it unset unless a run is memory-constrained. * **More devices is often slower.** See the next section; this matters more in MBIRTorch than the equivalent advice does in MBIRJAX. diff --git a/mbirtorch/_memory_ledger.py b/mbirtorch/_memory_ledger.py index 84f5cf5..df96684 100644 --- a/mbirtorch/_memory_ledger.py +++ b/mbirtorch/_memory_ledger.py @@ -28,10 +28,13 @@ ``torch.cuda.max_memory_allocated`` at the end of the reconstruction. That mode owns the peak counter (it resets it), so it is never on by default. -Two consumers share ONE per-view cost model. The projection drivers use +Two consumers share ONE view batch. The projection drivers use ``Projectors.view_batch_charge`` to choose a view batch; the ledger calls -the same function to price that batch's residency. The charge excludes the -call-fixed outputs by contract, so the ledger adds those itself, per phase. +the same function so it prices the batch the driver would actually run. The +charge excludes the call-fixed outputs by contract, so the ledger adds those +itself, per phase. It also reprices the batch when the body is a torch body, +because there the driver's number is a nominal slab used to bound the batch +and not a statement of what the batch holds; see TORCH_BODY_VIEW_SLABS. """ import math @@ -68,6 +71,19 @@ DIRECTION_CYLINDERS = 7 # apply_worker holds the direction and the scaled direction. APPLY_CYLINDERS = 2 +# How many gathered column cylinders a forward on the column-gather path +# holds at once (TomographyModel._sparse_forward_project_columns). The driver +# issues each batch's gather one batch ahead of the projection that reads it, +# so at the widest instant -- inside the gather that runs ahead -- a device +# holds three: the cylinder the projection is about to read, the pieces +# arriving from the slice-owners for the batch after it, and the concatenation +# those pieces are assembled into. +# +# The last batch of a pass has nothing to gather ahead of it, and a pass that +# fits in one batch never gathers ahead at all, so both hold two rather than +# three. The charge covers the widest instant, which is the rule the ledger +# keeps: it may charge more than a run needs but never less. +COLUMN_GATHER_RESIDENTS = 3 # Library workspace that torch allocates through its own caching allocator, # and that the ledger's array enumeration therefore cannot see. Measured as @@ -87,10 +103,48 @@ # configure_devices: the count is not searched and is never reduced, while the # empty-shard validation and the preflight still apply. DEVICE_COUNT_ENV_VAR = 'MBIRTORCH_NUM_DEVICES' +# How many per-view slabs one view batch of a TORCH BODY holds. A torch body +# is a projection body written as general torch code, which is what a geometry +# with no hand-written kernel runs. A hand-written kernel body declares what +# one of its views costs; a torch body declares nothing, so the driver prices +# it at ONE nominal slab -- (view batch, pixels, columns) floats -- and that +# single slab is what the ledger used to charge. +# +# A torch body holds a whole loop of those slabs at once. It walks the +# interpolation kernel one offset at a time, and each offset materializes an +# integer index array, a weight array and a gathered array of the slab's +# shape, none of which fuse away; the running output and the mapped centers +# stay live across the whole loop beside them. +# +# Measured 2026-08-10 on four H100s (job mg8), over the two geometries with no +# hand-written kernels, four problem sizes, and one, two and four devices. +# The runs whose measured peak is set by the projection itself need 12.9 +# slabs to cover it, and no run needs more. Charged at 14, eight percent +# above the tightest of those readings. +# +# ONE count covers both projection directions and both geometries, because the +# ledger cannot tell which body it holds: it sees only that the body declares +# no cost. The count is a measured multiplier and not a count of named +# arrays: the two geometries plainly do not hold the same number of slabs -- +# the runs of one need at most 7.0 where the other needs 12.9 -- and nothing +# in the plan distinguishes them, so the larger has to be charged to both. +# TORCH_BODY_CALIBRATION_BAND says what that costs the smaller. +TORCH_BODY_VIEW_SLABS = 14 + # The band the modeled peak must land in against the measured peak. The # lower bound is the one that matters: a ledger that under-predicts would let # a doomed run start, which is the failure this module exists to prevent. CALIBRATION_BAND = (1.00, 1.30) +# The same band for a reconstruction whose projection bodies are torch bodies. +# It is far wider than the one above for two reasons, both measured rather +# than assumed. One slab count has to cover two geometries that hold +# different numbers of slabs, since nothing in the plan distinguishes them. +# And two of the measured two-device runs peaked twice as high on one device +# as on the other from identical shards, which a per-device model built from +# shapes alone cannot reproduce: it must cover the higher device, so it +# over-charges the lower one by that factor. The widest over-charge measured +# is 5.74x, on the lower device of one of those two runs. +TORCH_BODY_CALIBRATION_BAND = (1.00, 5.80) class MemoryPreflightError(RuntimeError): @@ -184,11 +238,23 @@ class LedgerPlan: # ── knobs and model choices ────────────────────────────────────────────── forward_band: int = None back_band: int = None + # The pixel-column batch the forward's column gather assembles at once, + # or None when the forward walks slice bands instead. One field rather + # than a flag and a width, so the two can never disagree, and resolved by + # the model in plan_from_model rather than re-derived here. + column_pixel_batch: int = None qggmrf_cylinders: int = QGGMRF_CYLINDERS_COMPILED # (direction, num_pixels, band_cols) -> (view_batch, bytes_per_view), with # direction in {'forward', 'back'}. Defaults to a no-charge model so a # hand-built plan can exercise the state terms alone. view_charge: object = None + # Which of 'forward' and 'back' bind a torch body -- a body that declares + # no per-view cost of its own, so the ledger prices its views itself (see + # TORCH_BODY_VIEW_SLABS). The two directions are named separately because + # a model may bind a hand-written kernel one way and a torch body the + # other. Empty means both directions declare their own cost, which is + # what a hand-built plan gets: its charge reads exactly as before. + torch_body_directions: tuple = () @property def n_devices(self): @@ -250,38 +316,114 @@ def back_cols(i): return (plan.band_length(i, 'back') if plan.rows_track_slices else num_rows_dev) + def column_gather_slices(): + """The slice extent one column-gather call is handed: the WHOLE + device-form slice axis, padded tail included, because the gathered + cylinder spans every slice-owner at once.""" + return sum(int(block[0]) for block in plan.slice_blocks) + + def forward_call_pixels(num_pixels): + """How many pixel columns ONE forward call is handed: every pixel of + the pass by default, and one column batch on the column-gather path, + which is what makes that path's per-call terms fall.""" + if plan.column_pixel_batch: + return min(int(num_pixels), int(plan.column_pixel_batch)) + return int(num_pixels) + def forward_cols(i): """The forward call's band_cols: its voxel columns.""" - return (int(plan.recon_shape[2]) if n == 1 - else plan.band_length(i, 'forward')) + if n == 1: + return int(plan.recon_shape[2]) + if plan.column_pixel_batch: + return column_gather_slices() + return plan.band_length(i, 'forward') + + def band_slices(i, direction): + """The slice extent one projection call is handed: the whole slice + axis at one device, this owner's slice band under sharding, and the + whole device-form axis again on the column-gather path.""" + if n == 1: + return int(plan.recon_shape[2]) + if direction == 'forward' and plan.column_pixel_batch: + return column_gather_slices() + return plan.band_length(i, direction) + + def torch_body_batch(i, direction, num_pixels): + """What one view batch of a TORCH BODY holds. + + The body sweeps two axes -- the detector rows and the slice band it + was handed -- and every array in its interpolation loop spans the + view batch, the pixels, and whichever of those two axes is wider. + It holds TORCH_BODY_VIEW_SLABS of them at once, where the driver's + nominal charge prices one. + + The view batch itself stays the driver's own choice: only what that + batch is charged changes here, so the ledger and the driver still + agree on how many views one body call takes. + """ + if plan.view_charge is None: + return 0 + cols = back_cols(i) if direction == 'back' else forward_cols(i) + view_batch = int(plan.view_charge(direction, int(num_pixels), cols)[0]) + width = max(int(plan.sino_rows), int(band_slices(i, direction))) + return (TORCH_BODY_VIEW_SLABS * view_batch * int(num_pixels) + * width * _F32_BYTES) def back_batch(i, num_pixels): if not is_view_owner(i): return 0 + if 'back' in plan.torch_body_directions: + return torch_body_batch(i, 'back', num_pixels) return plan.batch_bytes('back', num_pixels, back_cols(i)) def forward_batch(i, num_pixels): if not is_view_owner(i): return 0 - return plan.batch_bytes('forward', num_pixels, forward_cols(i)) + # A call's own pixel count, which is the pass's on the banded path + # and one column batch on the column-gather path. + call_pixels = forward_call_pixels(num_pixels) + if 'forward' in plan.torch_body_directions: + return torch_body_batch(i, 'forward', call_pixels) + return plan.batch_bytes('forward', call_pixels, forward_cols(i)) def band_reduce(i, num_pixels): """The back reduce's co-residency on a slice-owner. - ``sum_band_to_owner`` moves ALL n partials onto the owner before the - summation loop begins, so the owner holds n arrays of one band plus - the running total. At three devices and above the old and the new - total coexist during a rebind, so the count is n + 2 there. Because - one band is the whole shard by default, this term is very nearly - INDEPENDENT of the device count: it reads 1.5x a full-volume cylinder - set at both two and four devices. Adding devices shrinks the - persistent set and leaves this where it was, which is why the - slice-band knob is the remedy the error message names for it. + ``sum_band_to_owner`` streams: it forms the running total for a band + once on the owner, then adds each arriving partial one row slab at a + time and frees the slab before the next one arrives. At the widest + instant the owner holds + + * the bands of its shard it has already reduced this pass and is + holding for the concatenation, at most ``shard - band`` slices, + * the running total for the band it is on, one band, + * the partial it produced itself, which the driver keeps alive + across the reduce, one band, + * one slab per arriving partial, each bounded by + ``_sharding.REDUCE_SLAB_BYTES``. + + That is ``shard + band`` slices of cylinder plus a bounded slab term, + which at the default band -- the whole shard -- is TWO + cylinder-shards. So it now falls as 1/n with the device count. The + old materialize-then-sum form held n whole bands plus the running + totals, which is the same number of bytes at every device count: it + measured 1.5x a full-volume cylinder set at both two and four + devices, and adding devices did not move it. + + The slab term does not shrink with the device count, but it is a + fixed number of bytes rather than a share of the volume. When a band + is smaller than one slab the whole band moves in one piece, which is + what the reduce always did, and this reads as the n + 1 bands that + then really are live. """ if n == 1 or not is_slice_owner(i): return 0 - copies = n + 1 if n == 2 else n + 2 - return copies * int(num_pixels) * plan.band_length(i, 'back') * _F32_BYTES + band = plan.band_length(i, 'back') + shard = plan.slice_blocks[i][0] + row_bytes = int(band) * _F32_BYTES + slab_rows = _sharding.reduce_slab_rows(int(num_pixels), row_bytes) + return (int(num_pixels) * (int(shard) + int(band)) * _F32_BYTES + + (n - 1) * slab_rows * row_bytes) def back_view_batches(i, num_pixels): """How many batches one worker's view loop runs, or None when this @@ -363,7 +505,15 @@ def forward_fixed(i): per-band pieces AND their concatenation (a row-aligned geometry, one whose detector row r comes from recon slice r), or the running partial AND the incoming one (a two-fan geometry such as cone, where one slice - projects onto many detector rows), so it pays twice.""" + projects onto many detector rows), so it pays twice. + + The COLUMN-GATHER forward holds one rather than two: its batches add + into the owner's block from inside the projector's view loop, so there + is no separate incoming block to hold beside it. The charge stays at + two anyway. It is shared with the banded path, which really does hold + both, and the ledger's rule is that it may charge more than a run needs + but never less -- so the column-gather path is deliberately over-charged + by one block here rather than given a term of its own.""" if not is_view_owner(i): return 0 return sino_dev(i) if n == 1 else 2 * sino_dev(i) @@ -377,11 +527,36 @@ def forward_band_copy(i, num_pixels): the copy is a full cylinder-shard on each device, on top of the device's own shard. Without this term the model falls below the measured peak on a large cone reconstruction at four devices. + + The column-gather path broadcasts no band at all, so this term is + zero there and ``forward_column_cylinder`` charges what it holds + instead. """ - if n == 1 or not is_view_owner(i): + if n == 1 or not is_view_owner(i) or plan.column_pixel_batch: return 0 return cyl(i, num_pixels) + def forward_column_cylinder(i, num_pixels): + """The gathered cylinder the column-gather forward assembles. + + ``_sharding.gather_column_band`` moves one batch of pixel columns + from every slice-owner and concatenates them, so what a view-owner + holds is that batch by the WHOLE device-form slice axis -- and, + unlike the band copy it replaces, that does not grow with the shard, + so it does not grow with the problem at a fixed batch. Three are live + at the widest instant, because the driver gathers one batch ahead of + the projection that reads it; see COLUMN_GATHER_RESIDENTS for which + three. + + Measured 2026-08-10 on four H100s, job mg10: ONE such cylinder read + 7.9, 15.8 and 31.5 MiB at batches 2048, 4096 and 8192 at 1008 slices, + which is the closed form exactly. + """ + if n == 1 or not is_view_owner(i) or not plan.column_pixel_batch: + return 0 + return (COLUMN_GATHER_RESIDENTS * forward_call_pixels(num_pixels) + * column_gather_slices() * _F32_BYTES) + def forward_view_batches(i, num_pixels): """How many batches one owner's forward view loop runs, or None when this plan prices no batch (a hand-built plan with no cost model). @@ -391,7 +566,7 @@ def forward_view_batches(i, num_pixels): if real_views <= 0 or plan.view_charge is None: return None view_batch = int(plan.view_charge( - 'forward', int(num_pixels), forward_cols(i))[0]) + 'forward', forward_call_pixels(num_pixels), forward_cols(i))[0]) return max(1, -(-int(real_views) // max(1, view_batch))) def forward_block_rows(i): @@ -414,33 +589,44 @@ def forward_block(i, num_pixels): """The view block the loop holds BESIDES the one the batch prices. ``Projectors.sparse_forward_project_view_range`` is ``block = - fwd_body(...)`` then ``out[...] = block``, with no release: python - evaluates the next call before it rebinds ``block``, so the loop holds - the outgoing block and the incoming one -- ``min(2, view_batches)`` - blocks. The back loop would hold the same two if it did not release - its block explicitly. - - ONE of those two is already inside ``forward batch``. A forward - body's output plane scales with the view batch, so each body's - ``_view_batch_cost`` charges it per view and says so; the back body's - cost model does not, its output being call-fixed at any batch. This - term is therefore the REMAINDER -- one block while the loop runs more - than a single batch, and nothing when it runs one, which is the whole - live set there. + fwd_body(...)`` then ``out[...] = block`` (or ``out[...].add_(block)`` + when the caller accumulates), with no release: python evaluates the next + call before it rebinds ``block``, so the loop holds the outgoing block + and the incoming one -- ``min(2, view_batches)`` blocks. Which of the + two arms runs does not change that count. The back loop would hold the + same two if it did not release its block explicitly. + + ONE of those two is already inside ``forward batch`` when the body + declares its own cost. A forward kernel body's output plane scales + with the view batch, so its ``_view_batch_cost`` charges it per view + and says so; the back body's cost model does not, its output being + call-fixed at any batch. Against a declared cost this term is + therefore the REMAINDER -- one block while the loop runs more than a + single batch, and nothing when it runs one, which is the whole live + set there. + + A TORCH BODY declares nothing, and what the ledger charges for it in + its place is the body's INTERNAL slab set, which does not include the + output plane. Nothing is already paid for there, so both blocks are + charged. The batch follows the pixel count of THIS call, so the subset phases - must pass their own subset size rather than the full index count. + must pass their own subset size rather than the full index count -- + and on the column-gather path a call's pixel count is one column + batch, which raises the view batch and with it this block. """ if not is_view_owner(i): return 0 batches = forward_view_batches(i, num_pixels) live = 2 if batches is None else min(2, batches) + already_paid = 0 if 'forward' in plan.torch_body_directions else 1 view_batch = 1 if plan.view_charge is not None: - view_batch = plan.view_charge('forward', num_pixels, + view_batch = plan.view_charge('forward', + forward_call_pixels(num_pixels), forward_cols(i))[0] - return ((live - 1) * int(view_batch) * forward_block_rows(i) - * num_channels * _F32_BYTES) + return ((live - already_paid) * int(view_batch) + * forward_block_rows(i) * num_channels * _F32_BYTES) # ── the persistent set ─────────────────────────────────────────────────── # One sinogram-shaped weights term, never two: when the caller supplies @@ -573,6 +759,8 @@ def back_phases(name, resident_terms, num_pixels, base, base_terms): ('init recon', per_dev(recon_dev)), ('voxel gather', per_dev(lambda i: cyl(i, p_full))), ('broadcast band', per_dev(lambda i: forward_band_copy(i, p_full))), + ('column cylinder', per_dev( + lambda i: forward_column_cylinder(i, p_full))), ('forward output', per_dev(forward_fixed)), ('forward block', per_dev(lambda i: forward_block(i, p_full))), ('forward batch', per_dev(lambda i: forward_batch(i, p_full))), @@ -701,6 +889,8 @@ def back_phases(name, resident_terms, num_pixels, base, base_terms): lambda i: sino_dev(i) if n > 1 and is_view_owner(i) else 0)), ('broadcast band', per_dev( lambda i: forward_band_copy(i, p_sub))), + ('column cylinder', per_dev( + lambda i: forward_column_cylinder(i, p_sub))), ('forward block', per_dev(lambda i: forward_block(i, p_sub))), ('forward batch', per_dev(lambda i: forward_batch(i, p_sub))), ], @@ -820,11 +1010,34 @@ def plan_from_model(model, devices, partition_sequence=None, weights=None, hessian_masked=model.get_params('use_ror_mask') is not False, forward_band=getattr(model, 'forward_project_slice_band', None), back_band=getattr(model, 'back_project_slice_band', None), + # Both read from the model's own resolvers rather than re-derived + # here: a charge that re-implements a driver rule is a charge that + # can be left behind when the rule moves. + column_pixel_batch=(model._forward_pixel_batch() + if model._column_gather_forward() else None), qggmrf_cylinders=qggmrf_cylinder_count(model), view_charge=charge, + torch_body_directions=torch_body_directions(model), ) +def torch_body_directions(model): + """Which projection directions this model runs as a torch body. + + A hand-written kernel body carries a ``_view_batch_cost`` attribute + stating what one of its views holds; general torch code carries nothing, + and the ledger prices those views itself (see TORCH_BODY_VIEW_SLABS). + The two directions are asked separately, because a model may bind a + kernel one way and a torch body the other, and because a kernel that is + unavailable on this machine falls back to the torch body it replaced -- + the charge has to follow the body that will actually run. + """ + fwd_body, back_body = model._view_batch_bodies() + return tuple(name for name, body in (('forward', fwd_body), + ('back', back_body)) + if getattr(body, '_view_batch_cost', None) is None) + + def _model_view_charge(model, n_devices): """A ``(direction, P, cols) -> (batch, bytes_per_view)`` closure over the bodies this model would actually bind.""" @@ -974,9 +1187,11 @@ def format_shortfall(ledger, rows, num_devices_tried, closest_count=None, + closest, '', 'Remedies, most effective first:', ' model.back_project_slice_band = ' - '# the band reduce barely shrinks with more', + '# shrinks every back projection transient', + ' ' + '# that is sized by a band, on top of what', ' ' - '# devices; this is its only lever', + '# more devices already save', ' model.view_batch_size = ' '# caps the projector batch transient', ' model.set_params(granularity=[...]) ' @@ -1046,8 +1261,11 @@ def calibration_report(ledger, devices): return rows -def format_calibration(rows): - low, high = CALIBRATION_BAND +def format_calibration(rows, band=None): + """The calibration table. ``band`` defaults to CALIBRATION_BAND; a + reconstruction whose projection bodies are torch bodies is judged against + TORCH_BODY_CALIBRATION_BAND instead.""" + low, high = band or CALIBRATION_BAND lines = ['memory ledger calibration (this mode owns ' 'torch.cuda.max_memory_allocated)', f'{"device":>10}{"modeled":>14}{"measured":>14}' diff --git a/mbirtorch/_sharding.py b/mbirtorch/_sharding.py index 21b6beb..74506cf 100644 --- a/mbirtorch/_sharding.py +++ b/mbirtorch/_sharding.py @@ -15,15 +15,29 @@ so the n=1 reconstruction path is unchanged. Under view/slice sharding the only data that crosses the recon<->sino -boundary is voxel-cylinder slice-bands (the sinogram is written locally on -its view-shard and never moves). That crossing is the banded adjoint pair: +boundary is voxel cylinders (the sinogram is written locally on its +view-shard and never moves). Two shapes of that crossing exist, and they +differ in which axis of the cylinder is cut. The banded adjoint pair cuts +the SLICE axis: - ``broadcast_band_to_views`` (forward / all-gather): copy a slice-band from its slice-owner to every view-owner. - ``sum_band_to_owner`` (back / reduce-scatter): sum each view-owner's band partials onto the band's slice-owner. + +``gather_column_band`` cuts the PIXEL axis instead: it assembles one batch of +pixel columns at every slice on one view-owner. A geometry whose slices +project onto a range of detector rows needs the whole slice axis before it +can produce any of its own rows, so a slice band buys it nothing, and the +forward driver gathers columns for it when that path is switched on. A +row-aligned geometry can produce its rows from a band and takes the same +gather anyway, because its kernel is markedly faster on the wider block of +values. Only the forward has the second shape; the back projection reduces +through ``sum_band_to_owner`` either way. """ +import contextlib +import threading import warnings from concurrent.futures import ThreadPoolExecutor @@ -234,6 +248,36 @@ def move_shard(x, target, dev2dev_safe=True): return torch.as_tensor(x.detach().cpu().numpy()).to(target) +#: How many bytes of one arriving partial the reduce moves at a time. A +#: REASONED default, not a measured knee -- the cluster measurement of this +#: change comes after it. The slab has to be large enough that the fixed +#: cost of one step (a python call, one device-to-device copy, one add) stays +#: small beside the step's own work: at 64 MiB the copy and the add take +#: hundreds of microseconds on any device-to-device link, against tens of +#: microseconds of launch and dispatch, so the host stays well ahead of the +#: devices and the streaming overhead is a few percent of the reduce. And it +#: has to be small enough to be negligible beside the band it streams: a +#: production band is gigabytes, so the slab is well under one percent of it. +#: A band smaller than one slab moves in a single piece, which is exactly +#: what the reduce did before, so nothing changes at small sizes. +REDUCE_SLAB_BYTES = 64 * 2 ** 20 + + +def reduce_slab_rows(num_rows, row_bytes): + """How many rows of a band partial :func:`sum_band_to_owner` moves per + step, given the bytes in one row of it. + + Shared with the memory ledger, which prices the transient this bounds: + the size the code moves and the size the model charges must not be able + to drift apart. + """ + # Never zero: the answer is a loop step, and a step of zero is an error + # even where the range it walks is empty. + if row_bytes <= 0: + return max(1, int(num_rows)) + return max(1, min(int(num_rows), int(REDUCE_SLAB_BYTES) // int(row_bytes))) + + def sum_band_to_owner(partials, owner, dev2dev_safe=True): """Move per-device partials onto ``owner`` and sum them there. @@ -241,11 +285,62 @@ def sum_band_to_owner(partials, owner, dev2dev_safe=True): view-sharding each device computed only a partial back projection (its own views' contribution) for some band of slices; the true value is the sum over devices, formed and left resident on the band's slice-owner. + + The sum is STREAMED in row slabs, and that is what bounds the owner's + peak. Moving every partial across first and then summing them held n + whole bands on the owner at once, and because one band is the whole shard + by default, that transient did not shrink as devices were added: n + devices each holding a band of 1/n of the volume is the same number of + bytes at every device count. Streaming leaves the owner holding its + running total and one bounded slab per source instead, so what it holds + ABOVE the total is a fixed number of bytes rather than a share of the + volume. + + The summation order is unchanged. Every element is still accumulated in + the order the partials are given, so the result is bit for bit what the + unstreamed reduce produced: streaming partitions the elements, and no + element's own sequence of additions is touched. + + The partials are read and never written. The first one is copied (or + moved) to make the running total, so a caller may still use its arrays + after the call. + + Args: + partials (list of tensor): one band partial per contributing device, + all of the same shape, summed in the order given. + owner (torch.device): the band's slice-owner, where the sum is formed + and left resident. + dev2dev_safe (bool): forwarded to :func:`move_shard`. """ - contribs = [move_shard(p, owner, dev2dev_safe=dev2dev_safe) for p in partials] - total = contribs[0] - for c in contribs[1:]: - total = total + c + if len(partials) == 1: + return move_shard(partials[0], owner, dev2dev_safe=dev2dev_safe) + total = move_shard(partials[0], owner, dev2dev_safe=dev2dev_safe) + if total is partials[0]: + # The first partial already lives on the owner, so move_shard handed + # back the caller's own tensor. Accumulate into a copy of it rather + # than writing through to an array the caller still holds. + total = total.clone() + num_rows = int(total.shape[0]) + row_bytes = (total.numel() // max(1, num_rows)) * total.element_size() + step = reduce_slab_rows(num_rows, row_bytes) + for start in range(0, num_rows, step): + stop = min(start + step, num_rows) + # Rows, not slices: a partial is (pixels, slices) with the slices + # contiguous, so a block of ROWS is a contiguous piece and each + # transfer stays a single flat copy. Every source's transfer for + # this slab is issued BEFORE any of them is consumed, so copies from + # different devices still overlap each other the way they did when + # whole bands were moved up front. + slabs = [move_shard(p[start:stop], owner, dev2dev_safe=dev2dev_safe) + for p in partials[1:]] + rows = total[start:stop] + for slab in slabs: + rows.add_(slab) + # Released here: the next iteration's list comprehension is evaluated + # BEFORE `slabs` is rebound, so without this the previous slabs stay + # live on the owner through the next slab's transfers, doubling the + # very transient this loop exists to bound. + slabs = None return total @@ -263,6 +358,191 @@ def broadcast_band_to_views(band, view_owners, dev2dev_safe=True): for dev in view_owners} +def gather_column_band(shard_tensors, p0, p1, target, dev2dev_safe=True): + """Gather one batch of pixel columns, at EVERY slice, onto ``target``. + + The forward's second transfer primitive, built from :func:`move_shard` + exactly as :func:`broadcast_band_to_views` is. Each slice-owner holds + the same pixel columns for its own slices, so moving every owner's + ``[p0:p1]`` rows to one device and concatenating them along the slice + axis assembles those columns' whole cylinder there. + + This is the cross-device shape a geometry needs when one recon slice + projects onto a RANGE of detector rows: such a view-owner cannot produce + any of its own rows from a slice band, because every slice contributes to + the rows it owns. It takes a narrow column of pixels at every slice + instead. What one gather costs is then set by the width of the column + batch and not by the device count, which is what makes the shape usable + at volumes where a whole assembled cylinder would not fit. A row-aligned + geometry, which could work from a band, takes the same gather for a + performance reason instead: what it gets back is a full-width block of + values, which is the width regime its kernel is efficient in. + + The concatenation is in shard order, which is global slice order, and it + keeps the device form's padded slice tail rather than trimming it. The + tail is held at zero by the model, a zero voxel contributes nothing + through a projection, and the geometry bodies anchor their z geometry on + the real slice count from the params rather than on the width of the + array they are handed -- so the tail is inert, and trimming it would only + force a non-contiguous copy inside the projector. + + This changes which device assembles which voxels, never which device + produces which sinogram rows, so it has no adjoint of its own: the back + projection is untouched and still reduces through + :func:`sum_band_to_owner`. + + Args: + shard_tensors (sequence of tensor): the slice-sharded cylinders, each + (num_pixels, local_slices), in global slice order. + p0 (int): first pixel column of the batch. + p1 (int): one past the last pixel column of the batch. + target (torch.device): the view-owner the cylinder is assembled on. + dev2dev_safe (bool): forwarded to :func:`move_shard`. + + Returns: + tensor: (p1 - p0, total_slices) on ``target``. + """ + pieces = [move_shard(t[p0:p1], target, dev2dev_safe=dev2dev_safe) + for t in shard_tensors] + return pieces[0] if len(pieces) == 1 else torch.cat(pieces, dim=1) + + +# ── copy streams for the column gather (CUDA only) ─────────────────────────── +# One extra CUDA stream per device, used for nothing but the column gather's +# cross-device copies. A stream runs its work in the order it was given, one +# item at a time, so copies left on the stream a device projects on can only +# take turns with those projections however early they are issued -- torch +# issues a cross-device copy on the SOURCE device's current stream and orders +# the DESTINATION device's current stream behind it, and for the gather's +# worker threads both of those are the default stream the device projects on. +# A stream of their own is what lets a copy and a projection run at once. +# +# Cached per device index and created once, the way projectors.py caches its +# compiled bodies: the lock is taken only to CREATE a stream, so the worker +# threads that ask for one every batch find it already there and stay +# lock-free. +_COPY_STREAMS = {} +_COPY_STREAM_LOCK = threading.Lock() + + +def copy_stream(device): + """The dedicated copy stream for ``device``, or None when it has none. + + None is returned for every non-CUDA device, and it is the signal the + callers below read as "this device has no streams to arrange": each of + them then does the plain synchronous thing, which is what the CPU and MPS + paths have always done. + """ + device = torch.device(device) + if device.type != 'cuda': + return None + index = (device.index if device.index is not None + else torch.cuda.current_device()) + stream = _COPY_STREAMS.get(index) + if stream is None: + with _COPY_STREAM_LOCK: + stream = _COPY_STREAMS.get(index) + if stream is None: + stream = torch.cuda.Stream(device=index) + _COPY_STREAMS[index] = stream + return stream + + +def _gather_stream_devices(shard_tensors, target): + """The distinct CUDA devices one gather touches: every shard's device and + the target it assembles on. Ordered by device index so that the nested + stream contexts are always entered in the same order.""" + seen = {} + for dev in [t.device for t in shard_tensors] + [torch.device(target)]: + if dev.type == 'cuda': + index = (dev.index if dev.index is not None + else torch.cuda.current_device()) + seen[index] = torch.device('cuda', index) + return [seen[index] for index in sorted(seen)] + + +def open_copy_streams(devices): + """Let the copy streams start: each waits for its device's compute stream. + + The shards a gather reads were written by earlier kernels on the compute + stream, and a copy stream knows nothing of that stream's ordering, so + without this a copy could read a shard before the kernel that filled it + had finished. Called once per forward rather than per batch: it orders + the copy stream behind everything queued so far, which covers every batch + that follows. + """ + for dev in devices: + stream = copy_stream(dev) + if stream is not None: + stream.wait_stream(torch.cuda.current_stream(torch.device(dev))) + + +def close_copy_streams(devices): + """The other half of :func:`open_copy_streams`: each compute stream waits + for its copy stream. + + A copy READS a slice-owner's shard, and whatever writes that shard next + runs on the compute stream. Nothing else orders those two, so without + this a later update could overwrite a shard while a copy was still + reading it. + """ + for dev in devices: + stream = copy_stream(dev) + if stream is not None: + torch.cuda.current_stream(torch.device(dev)).wait_stream(stream) + + +def gather_column_band_async(shard_tensors, p0, p1, target, dev2dev_safe=True): + """:func:`gather_column_band`, issued on the copy streams. + + The values are the same either way; what this adds is that the copies do + not go into the queue the projections run in, so a gather can be moving + while an earlier batch is projected. + + Returns: + (tensor, ready): the assembled cylinder, and an event that fires once + its copies have landed -- or None for the event off CUDA, where the + copies are already finished by the time this returns. + """ + stream = copy_stream(target) + if stream is None: + return gather_column_band(shard_tensors, p0, p1, target, + dev2dev_safe), None + # BOTH ends of every copy have to be on a copy stream: torch issues the + # copy on the source's current stream and orders the destination's current + # stream behind it, so leaving either end on its default stream would put + # the copy straight back in the queue the projections run in. + with contextlib.ExitStack() as stack: + for dev in _gather_stream_devices(shard_tensors, target): + stack.enter_context(torch.cuda.stream(copy_stream(dev))) + cylinder = gather_column_band(shard_tensors, p0, p1, target, + dev2dev_safe) + ready = torch.cuda.Event() + ready.record(stream) + # The cylinder was allocated on the copy stream and is read on the compute + # stream. Without this the caching allocator would be free to hand its + # block to the next gather the moment python drops the name, while the + # projection was still reading it. This covers the arriving pieces too: + # they are allocated and concatenated on the one copy stream, and the only + # one that ever escapes is the single-shard case, where the piece IS the + # cylinder returned here. + cylinder.record_stream(torch.cuda.current_stream(torch.device(target))) + return cylinder, ready + + +def wait_for_column_band(target, ready): + """Hold ``target``'s compute stream until one batch's copies have landed. + + The event is per batch and is waited on immediately before the projection + that reads that batch. Waiting on the copy stream as a whole instead + would also wait for the batch gathered ahead, which is exactly the work + meant to be moving during this projection, and the overlap would collapse + back into taking turns. + """ + if ready is not None: + torch.cuda.current_stream(torch.device(target)).wait_event(ready) + + # ── per-device threaded execution (the mbirjax thread_execution.py port) ────── def device_pool(n): """A reusable thread pool for repeated :func:`run_per_device` calls. diff --git a/mbirtorch/_widening_floors.py b/mbirtorch/_widening_floors.py index 8147e5c..22c3a32 100644 --- a/mbirtorch/_widening_floors.py +++ b/mbirtorch/_widening_floors.py @@ -122,52 +122,56 @@ ('parallel', 2): Floor( family='parallel', count=2, elements=88_080_384, cell=(512, 448, 384), against=1, - bracket=Bracket(losing_cell=(384, 336, 288), losing_speedup=0.64, - winning_cell=(512, 448, 384), winning_speedup=1.23), - spread=0.09623, gpu=MEASURED_GPU, config=MEASURED_CONFIG, - measured='2026-08-10', commit='a880d9c', + bracket=Bracket(losing_cell=(384, 336, 288), losing_speedup=0.80, + winning_cell=(512, 448, 384), winning_speedup=1.21), + spread=0.05155, gpu=MEASURED_GPU, config=MEASURED_CONFIG, + measured='2026-08-11', commit='4a222c7', largest_tested=297_271_296, - note='unchanged by the 2026-08-10 refresh. The spread comes from ' - 'the 384-class n=1 runs, the noisiest in the family; at the ' - 'floor shape itself n=2 wins by 1.23x'), + note='unchanged by the 2026-08-11 refresh, which re-measured ' + 'every row under the column-gather forward default. The ' + '384-class shape still loses at 0.80x, and the floor shape ' + 'wins by 1.21x'), ('parallel', 4): Floor( - family='parallel', count=4, elements=1_023_934_464, - cell=(1024, 1008, 992), against=2, - bracket=Bracket(losing_cell=(768, 672, 576), losing_speedup=0.74, - winning_cell=(1024, 1008, 992), winning_speedup=1.67), - spread=0.008199, gpu=MEASURED_GPU, config=MEASURED_CONFIG, - measured='2026-08-10', commit='a880d9c', + family='parallel', count=4, elements=297_271_296, + cell=(768, 672, 576), against=2, + bracket=Bracket(losing_cell=None, losing_speedup=None, + winning_cell=(768, 672, 576), winning_speedup=1.10), + spread=0.02116, gpu=MEASURED_GPU, config=MEASURED_CONFIG, + measured='2026-08-11', commit='4a222c7', largest_tested=1_023_934_464, - note='measured against n=2, which still wins at the 768-class ' - 'shape'), + note='the floor MOVED DOWN, from the 1024-class shape to the ' + '768-class, under the column-gather forward: four devices ' + 'now clear two by 1.10x at the 768-class shape and by 1.47x ' + 'at the 1024-class. No losing shape is recorded because no ' + 'smaller shape was tried once the 768-class won; if this ' + 'admission is wrong, the cost is bounded by the 1.10x ' + 'margin against its 2.1 percent spread'), ('cone', 2): Floor( family='cone', count=2, elements=88_080_384, cell=(512, 448, 384), against=1, - bracket=Bracket(losing_cell=None, losing_speedup=None, - winning_cell=(512, 448, 384), winning_speedup=1.02), - spread=0.005233, gpu=MEASURED_GPU, config=MEASURED_CONFIG, - measured='2026-08-10', commit='a880d9c', - largest_tested=1_023_934_464, - note='the FIRST admission size ever measured for this entry: ' - 'before the 2026-08-10 refresh, no size had one. MARGINAL, ' - 'on a 1.02x win clearing a 0.52 percent spread, and with no ' - 'losing shape recorded because no smaller shape was tried ' - 'once the 512-class shape won. If this admission is wrong, ' - 'the cost is a few percent, by the measured asymmetry'), + bracket=Bracket(losing_cell=(384, 336, 288), losing_speedup=0.78, + winning_cell=(512, 448, 384), winning_speedup=1.21), + spread=0.02342, gpu=MEASURED_GPU, config=MEASURED_CONFIG, + measured='2026-08-11', commit='4a222c7', + largest_tested=297_271_296, + note='unchanged, and no longer marginal: the admission that ' + 'cleared by 1.02x on 2026-08-10 clears by 1.21x under the ' + 'column-gather forward, and this refresh recorded the ' + 'losing shape the first measurement never tried'), ('cone', 4): Floor( family='cone', count=4, elements=1_023_934_464, - cell=(1024, 1008, 992), against=1, - bracket=Bracket(losing_cell=(768, 672, 576), losing_speedup=0.98, - winning_cell=(1024, 1008, 992), winning_speedup=1.16), - spread=0.003432, gpu=MEASURED_GPU, config=MEASURED_CONFIG, - measured='2026-08-10', commit='a880d9c', + cell=(1024, 1008, 992), against=2, + bracket=Bracket(losing_cell=(768, 672, 576), losing_speedup=0.95, + winning_cell=(1024, 1008, 992), winning_speedup=1.45), + spread=0.01853, gpu=MEASURED_GPU, config=MEASURED_CONFIG, + measured='2026-08-11', commit='4a222c7', largest_tested=1_023_934_464, - note='the refresh narrowed the bracket from 512-to-1024 down to ' - '768-to-1024; the floor did not move. Measured against ' - 'n=1, because cone n=2 was not admitted anywhere when this ' - 'row was set. Now that cone n=2 has a floor, the crossover ' - 'rule means the next refresh re-derives this row against ' - 'n=2'), + note='the floor did not move, but this refresh re-derived it ' + 'against n=2, as the crossover rule requires now that cone ' + 'n=2 is admitted -- the change the previous entry ' + 'anticipated. At the floor shape n=4 clears n=2 by 1.45x, ' + 'and the 768-class shape sits just under admission at ' + '0.95x, so the bracket is tight'), } # ── the projection-cost inputs the floors were measured against ────────────── @@ -175,13 +179,20 @@ #: module-level chunk constants and the budget class attributes these files #: carry are exactly the kind of tuning that moves a crossover without #: touching any function this table names, so a function-level hash would -#: miss them. -COST_INPUT_FILES = ('triton_parallel.py', 'triton_cone.py', 'projectors.py') +#: miss them. ``_sharding.py`` is here because it holds the cross-device +#: transfer primitives the multi-device drivers are built from, and how much +#: those move is most of what a wider device count costs. +COST_INPUT_FILES = ('triton_parallel.py', 'triton_cone.py', 'projectors.py', + '_sharding.py') #: Methods of TomographyModel that drive the multi-device projections. The #: rest of that module moves for reasons unrelated to projection cost, so the -#: hash is taken over these two sources rather than the whole file. +#: hash is taken over these sources rather than the whole file. The column +#: gather is a third driver rather than a branch of the first, so it is named +#: here in its own right; leaving it out would let the pixel batch it walks +#: change without anything noticing. COST_INPUT_METHODS = ('_sparse_forward_project_sharded', + '_sparse_forward_project_columns', '_sparse_back_project_sharded') #: sha256 of each cost input as of the measurement above -- the recorded @@ -190,10 +201,14 @@ BLESSED_COST_HASHES = { 'TomographyModel._sparse_back_project_sharded': '8a39fb4d97a9573933520ce780eae5dd2097e5a068caa3ee2178114ba8989772', + 'TomographyModel._sparse_forward_project_columns': + '73f545dbd63188d6668a59d1707200a9cd065a0fbed3fcd929d713af77e01993', 'TomographyModel._sparse_forward_project_sharded': - 'f2a1fff6d1ea2627abfa4d02f1b5ad08e80383f4eef6bd037e4743c200b9f7b2', + '546201c90075a19f5ffe055c2becee6716417aa52e9e5f178885e7a68aae60f3', + '_sharding.py': + '424ada53243fa9f486cf139ee8564d21162ea791b9ef59ed949d0fa8a85d9b35', 'projectors.py': - '6977a6181accbaee9235246ce2cc59f17869d9666daae91a3748b2c21f143cf6', + '68e812790a963b92519169fe4a04e667c587ff70919ed574533e4c52c891698a', 'triton_cone.py': '8d3820c2101f8d3fbb7823f2d9b6e6e6253164bd14a2c276d167d9ba0a135154', 'triton_parallel.py': @@ -214,7 +229,7 @@ #: green the test leaves this behind, and the test says so. Recomputed and #: printed by ``refresh_widening_floors.py --bless``. TABLE_CHECKSUM = \ - 'aa728b2070772ef627874d3bfc11206088ee3666f8e240319e50bea777596886' + 'e68d1c7fb7e0d6a25e62f1053e562ce851651a12f70992c38fb7485e019613ae' # ── the env knob ───────────────────────────────────────────────────────────── diff --git a/mbirtorch/cone_beam.py b/mbirtorch/cone_beam.py index efc03d7..12116ee 100644 --- a/mbirtorch/cone_beam.py +++ b/mbirtorch/cone_beam.py @@ -343,6 +343,14 @@ def __init__(self, sinogram_shape, angles, source_detector_dist, source_iso_dist # higher than parallel's: its n=2 has no measured admission point at all. _floor_family = 'cone' + # Cone is the geometry the multi-device forward's column gather was + # measured on (see TomographyModel._column_gather_forward for what else + # has to hold before it runs, and _sparse_forward_project_columns for the + # numbers). The path runs by default since its speed, value, and memory + # gates passed (2026-08-11, four H100s); forward_column_gather = False + # restores the banded walk. + column_gather_geometry = True + def create_projectors(self): super().create_projectors() # Warm the DC-damping profile and its per-device compiled instances @@ -764,6 +772,12 @@ def fdk_recon(self, sinogram, filter_name="ramp", output_sharded=False): applies no short-scan redundancy weighting; for helical scans it is approximate regardless. Best used as an initializer for ``recon()``. """ + # Settle the device layout before the first large allocation, as + # recon() does: a no-op when the user already chose devices; + # otherwise the automatic selection runs here, so a bare FDK call + # spreads across the GPUs instead of landing whole on one (the A2 + # gap that failed the full-resolution MAR runs). + self._apply_device_policy() # Place once at entry so the filter receives device-form data (a no-op # when already placed; a single device is the trivial 1-shard case). # The pipeline then stays on-device throughout -- fdk_filter then diff --git a/mbirtorch/denoising.py b/mbirtorch/denoising.py index b1f80e8..bc9138f 100644 --- a/mbirtorch/denoising.py +++ b/mbirtorch/denoising.py @@ -439,8 +439,9 @@ def apply_worker(j, dev): ell1_accum = ell1_accum + combine_on_lead(ell1_parts) alpha_accum = alpha_accum + alpha - # The one host synchronization per pass: the convergence - # test and the two logged histories need Python numbers. + # The three host reads per pass, all at this one + # synchronization point: the convergence test and the two + # logged histories need Python numbers. image_l1 = combine_on_lead([torch.sum(torch.abs(t)) for t in flat_image.tensors]) nmae = float(ell1_accum) / float(image_l1) diff --git a/mbirtorch/parallel_beam.py b/mbirtorch/parallel_beam.py index ef29670..ac40245 100644 --- a/mbirtorch/parallel_beam.py +++ b/mbirtorch/parallel_beam.py @@ -157,6 +157,31 @@ def get_magnification(self): # automatic device count (see _widening_floors). _floor_family = 'parallel' + # Never call the compiled parallel bodies with a single pixel: on linux + # with torch 2.13.0, CPU inductor miscompiles that one-pixel case in both + # bodies and lands the pixel's footprint one detector channel off (6.56e-02 + # relative error on the forward, 5.04e-02 on the back; eager is right, and + # so is every width of two or more). The driver pads a one-pixel call to + # two and takes the padding back out, outside the compiled region -- + # projectors.forward_at_min_pixel_width holds the full measurement and the + # argument that the padding cannot change a value. Cone beam does not need + # this and does not declare it. + min_compiled_pixel_width = 2 + + # Parallel takes the multi-device forward's column gather for a reason of + # its own: the forward kernel runs about twice as efficiently per slice on + # a full-width block of values as on the shard-width blocks the banded + # walk hands it at more than one device, and the gather hands it full + # width whatever the device count (measured 2026-08-10 on one H100, at + # 0.0411 ms per slice on a 1008-wide block against 0.0823 on a 504-wide + # one with the device count held at one). Cone declares the same + # attribute because a slice band buys its kernel nothing at all; see + # TomographyModel._column_gather_forward for what else has to hold before + # the path runs. It runs by default since its speed, value, and memory + # gates passed (2026-08-11, four H100s); forward_column_gather = False + # restores the banded walk. + column_gather_geometry = True + def get_psf_radius(self): """Computes the integer radius of the PSF kernel for parallel beam projection: the maximum number of detector channels on either side of diff --git a/mbirtorch/projectors.py b/mbirtorch/projectors.py index 18edcf9..e331283 100644 --- a/mbirtorch/projectors.py +++ b/mbirtorch/projectors.py @@ -127,6 +127,91 @@ def guarded(*args, **kwargs): return guarded +# ── the minimum pixel width a compiled body is called at ───────────────────── +# A model may declare that its compiled bodies must not be called with fewer +# than N pixels (``min_compiled_pixel_width``, see TomographyModel). The two +# wrappers below pad a narrower call up to N and undo the padding on the way +# out. They are applied OUTSIDE torch.compile, around the callable +# maybe_compile returns, so the padding is ordinary python that dynamo never +# traces -- padding inside the body would be traced and specialized with it, +# which is exactly what has to be avoided. +# +# Why they exist (measured 2026-08-11, linux CPU, torch 2.13.0): inductor +# miscompiles the one-pixel specialization of both fused parallel-beam bodies. +# A one-pixel call puts that pixel's horizontal-fan footprint one whole +# detector channel away from where the same pixel lands in a call with more +# pixels -- 6.56e-02 relative error on the forward (the pixel's mass at +# channels {4, 5} instead of {3, 4}) and 5.04e-02 on the back, on a seeded +# 8x6x8 test cell. Eager is correct at one pixel (1.05e-07), every width of +# two or more is correct compiled, and a one-pixel call is correct once the +# process has compiled the body at a larger width, so what is wrong is the +# one-pixel compile itself. The cone bodies do not have the defect, and macOS +# inductor compiles the same one-pixel body correctly. One-pixel calls are +# ordinary: sparse_forward_project with a single index makes one, and the +# column gather's pixel batching makes one whenever a batch, or the remainder +# of a batch, is a single pixel. +# +# The driver's view-batch charge is computed from the REAL pixel count, before +# the padding: it prices the transient of a call this small at a batch far +# below any cap, so the one padded column cannot move it. + + +def _callable_name(fn, fallback): + """A readable name for a wrapped callable, for the wrapper's own name.""" + return getattr(fn, '__name__', fallback) + + +def forward_at_min_pixel_width(compiled, min_width): + """The forward body with narrow pixel batches padded to ``min_width``. + + The padded columns carry zero values at a repeated -- hence in-range -- + pixel index. The forward output has no pixel axis: the fan bins each + pixel's weighted row into the detector channels with index_add_, so a + zero-valued column adds exactly 0.0 wherever it lands and the padded call + returns bit-identical values with nothing to slice off (verified against + the eager body). + """ + def forward_padded(values, pixel_indices, *args, **kwargs): + width = int(pixel_indices.shape[0]) + if width == 0 or width >= min_width: + return compiled(values, pixel_indices, *args, **kwargs) + pad = min_width - width + wide_values = torch.cat( + [values, values.new_zeros((pad,) + tuple(values.shape[1:]))]) + wide_indices = torch.cat([pixel_indices, + pixel_indices[-1:].repeat(pad)]) + return compiled(wide_values, wide_indices, *args, **kwargs) + + forward_padded.__name__ = f'padded_{_callable_name(compiled, "forward")}' + return forward_padded + + +def back_at_min_pixel_width(compiled, min_width): + """The back body with narrow pixel batches padded to ``min_width``. + + The back output DOES carry the pixel axis, so here the padding repeats the + last real pixel index and the extra rows are sliced off again. Every + output row is computed from its own pixel alone (the fan gathers per pixel + and sums over views), so the rows that stay are the rows the narrow call + would have produced -- exactly, not to a tolerance (verified against the + eager body, at coeff_power 1 and 2). + """ + def back_padded(sino_batch, pixel_indices, *args, **kwargs): + width = int(pixel_indices.shape[0]) + if width == 0 or width >= min_width: + return compiled(sino_batch, pixel_indices, *args, **kwargs) + pad = min_width - width + wide_indices = torch.cat([pixel_indices, + pixel_indices[-1:].repeat(pad)]) + block = compiled(sino_batch, wide_indices, *args, **kwargs) + # Cloned rather than returned as a view, so the caller's output owns + # its memory and does not keep the padded block alive. + return block[:width].clone() + + back_padded.__name__ = f'padded_{_callable_name(compiled, "back")}' + return back_padded + + def compile_serialized(): """The process-wide compile lock, as a context manager -- for HAND-WRITTEN kernel paths only:: @@ -258,11 +343,30 @@ def __init__(self, model): fwd_body, back_body = model._view_batch_bodies() use_compile = model.compile_enabled n_dev = model.sino_placement.n_devices + min_width = int(getattr(model, 'min_compiled_pixel_width', 1)) + + def bind(body, pad_narrow, i): + """One device's bound body: compiled, then wrapped when the model + declares a minimum pixel width AND the binding really did compile. + + The identity test is the whole gate. maybe_compile hands back the + function itself when compilation is off and when the body is a + hand-written kernel (``_mbirtorch_no_compile``); neither can be + miscompiled, so neither needs the workaround, and leaving them + alone keeps the two things callers read off a bound body -- its + identity and its ``_view_batch_cost`` attribute -- exactly as they + were. Every driver, plain and sharded, reads its body from these + two lists, so this is the one place per direction to wrap.""" + bound = maybe_compile(body, use_compile, instance_key=i) + if min_width > 1 and bound is not body: + bound = pad_narrow(bound, min_width) + return bound + self._fwd_body_per_dev = [ - maybe_compile(fwd_body, use_compile, instance_key=i) + bind(fwd_body, forward_at_min_pixel_width, i) for i in range(n_dev)] self._back_body_per_dev = [ - maybe_compile(back_body, use_compile, instance_key=i) + bind(back_body, back_at_min_pixel_width, i) for i in range(n_dev)] # View parameters, read from the CURRENT params at every projector # build (create_projectors re-runs on reconfigure/recompile, closing @@ -351,7 +455,8 @@ def view_batch_charge(self, body, num_pixels, band_cols, args, def sparse_forward_project_view_range(self, band_values, pixel_indices, view_range, slice_start=0, - dev_index=0, plan=None): + dev_index=0, plan=None, + accumulate_into=None): """Forward-project voxel values into ONE view-owner's sinogram block: the single forward loop -- the single-device full-range form is the adapter below over (0, num_views). The geometry body owns all geometry, @@ -359,6 +464,17 @@ def sparse_forward_project_view_range(self, band_values, pixel_indices, transient budget, and assembly (output sized lazily from the first block, so the driver never derives geometry-specific shapes). + ``accumulate_into`` lets a caller that runs this loop repeatedly -- the + column-gather forward, once per pixel batch -- add straight into the + block it is building instead of receiving a fresh one to add itself. + That merges two full-block passes into one and drops one full-block + allocation per call; see the accumulation comment in + ``TomographyModel._sparse_forward_project_columns`` for why it is worth + doing and why the values do not move. The parameter is added HERE, on + a plain python method, and not to the geometry body: the bodies are + torch.compile'd per device with shape-keyed caches, so a new argument + there would recompile every one of them. + Args: band_values: (P, cols) voxel cylinders (or a slice band), on this owner's device. @@ -371,6 +487,9 @@ def sparse_forward_project_view_range(self, band_values, pixel_indices, dev_index (int): which per-device compiled instance to use. plan: the memoization slot for a future sorted/CSR stream variant (per pixel-subset x view-range); unused today. + accumulate_into: an existing block of the shape this call returns. + Given one, the loop ADDS into it and returns it; given None + (every other caller), it allocates the block and writes. Returns: (v1 - v0, rows_or_band, num_channels) on the input's device. @@ -382,7 +501,12 @@ def sparse_forward_project_view_range(self, band_values, pixel_indices, vb_size = self._effective_view_batch(fwd_body, pixel_indices.shape[0], band_values.shape[-1], args) view_params = self._view_params_per_dev[dev_index] - out = None + out = accumulate_into + # Whether this call adds into a block it was handed or fills one of its + # own, decided ONCE here rather than per view batch: an accumulating + # call adds every batch, including the first, because the block already + # holds earlier calls' work. + adding = out is not None for v in range(v0, v1, vb_size): view_params_batch = view_params[v:min(v + vb_size, v1)] block = fwd_body( @@ -391,7 +515,14 @@ def sparse_forward_project_view_range(self, band_values, pixel_indices, if out is None: out = torch.empty((v1 - v0,) + tuple(block.shape[1:]), dtype=block.dtype, device=block.device) - out[v - v0:v - v0 + block.shape[0]] = block + rows = slice(v - v0, v - v0 + block.shape[0]) + # View batches cover DISJOINT rows of the block, so neither arm + # sums anything across this loop -- assignment and addition touch + # each row exactly once either way. + if adding: + out[rows].add_(block) + else: + out[rows] = block return out def sparse_back_project_view_range(self, local_sino, pixel_indices, diff --git a/mbirtorch/tomography_model.py b/mbirtorch/tomography_model.py index 10e58cb..160f401 100644 --- a/mbirtorch/tomography_model.py +++ b/mbirtorch/tomography_model.py @@ -38,6 +38,30 @@ _F32_EPS = float(np.finfo(np.float32).eps) +# ── the multi-device forward's column gather ───────────────────────────────── +# How many pixel columns one gathered cylinder covers (see +# TomographyModel._forward_pixel_batch and _sparse_forward_project_columns). +# The cylinder is this many columns by the whole slice axis, so this is the +# knob that bounds the cross-device transient on that path. Measured +# 2026-08-10 on four H100s, job mg10: per-device forward time fell at every +# batch tried -- 2048, 4096, 8192 -- and was still falling at the largest. +# The sweep above it ran the next night (job mg11, same machines, 1K cells): +# 16384 and 32768 kept improving the composed wall by a further 4 to 15 +# percent depending on geometry and device count, so the knee is still not +# bracketed. 8192 stays the default anyway, because those readings come from +# a 1K harness and production runs at 2K and above, where the batch's +# transient grows with the slice axis and the sweep has not been run. Set +# forward_project_pixel_batch on the model to override. +FORWARD_PIXEL_BATCH = 8192 + +# Forces the column gather on ('1', 'true', 'yes', 'on') or off ('0', +# 'false', 'no', 'off') whatever the model attribute says. Read per call, +# like the other environment knobs, so one session can run both shapes -- the +# comparison the value gate for this path is read from. +COLUMN_GATHER_ENV_VAR = 'MBIRTORCH_FORWARD_COLUMN_GATHER' +_COLUMN_GATHER_ON_VALUES = ('1', 'true', 'yes', 'on') +_COLUMN_GATHER_OFF_VALUES = ('0', 'false', 'no', 'off') + # ── compiled updater glue (module level, one compile per process) ───────────── # Eagerly there were ~20 kernel launches per subset between the projector @@ -369,25 +393,83 @@ def _slice_band_length(slices_per_dev, n_dev, num_pixels, fixed_band=None): DEFAULT = one band per slice-owner (the whole shard). This differs from mbirjax deliberately, on measurement: mbirjax's sweeps found time flat across B, so it streams by default for the memory win, but - the torch banded pass pays a fixed orchestration cost per band - (eager fan-out), and splitting the shard into sub-bands was measured - on four H100s at 47 to 66 percent MORE reconstruction time at two - devices, for peak-memory savings of 0 to 61 percent -- far more - time than the memory is worth on this path. + the torch banded pass pays a fixed orchestration cost per band. + With the compiled kernels in place, splitting the shard into + sub-bands was measured on four H100s at 2 to 23 percent more busy + time at parallel 1024 with two devices, depending on the walk (job + mg10, 2026-08-10; an earlier pre-kernel reading of 47 to 66 percent + overstated the cost). The one exception, a 9.5 percent win at the + 63-slice walk, is non-monotonic and unexplained, and is not a basis + for a default. Time buys nothing back here because a single torch device never runs the banded drivers at all (the trivial fast path uses the plain projectors), so mbirjax's stream-even-at-n=1 rationale is void. A smaller B remains a real MEMORY lever (the per-band broadcast - copy, the per-band partial, and each slice-owner's reduce gather all - scale with B; measured n=4 @512: 6.6 to 2.6 GiB for +8 percent - time). Set ``forward_project_slice_band`` / - ``back_project_slice_band`` on the model to opt in with a fixed B - when a run is memory-constrained. Every result is capped at - slices_per_dev so a band never crosses a slice-owner boundary.""" + copy, the per-band partial, and the running total each slice-owner + reduces into all scale with B; the same mg10 sweep read per-device + peaks of 11.84 to 11.97 GB across the sub-band walks against 12.48 GB + at the default, with total copied bytes unchanged). That sweep + predates the streamed reduce, which took the default-B reduce from n + whole bands down to two plus a bounded slab, so expect a narrower gap + than those peaks show. Set + ``forward_project_slice_band`` / ``back_project_slice_band`` on the + model to opt in with a fixed B when a run is memory-constrained. + Every result is capped at slices_per_dev so a band never crosses a + slice-owner boundary.""" b = fixed_band if fixed_band else slices_per_dev return min(int(b), slices_per_dev) + def _column_gather_forward(self): + """Whether the multi-device forward gathers pixel COLUMNS instead of + walking slice bands. + + Three things have to agree, and each guards a different mistake. + + The GEOMETRY must be one the column gather has been measured for. + ``column_gather_geometry`` is set by cone beam and by parallel beam, + which want the same full slice range for two different measured + reasons. Cone NEEDS it: one slice projects onto a range of detector + rows, so a band-sized call still writes every row and costs what a + full call costs. Parallel merely wants it: its forward kernel runs + about twice as efficiently per slice on a full-width block of values + as on the shard-width blocks the banded walk hands it at more than + one device. Translation and multiaxis share cone's banded + branch and its band-independent per-call cost, so the shape should + help them too, but neither has ever been timed on it and neither + should be switched over on an argument. + + The SWITCH must not be off. ``forward_column_gather`` unset means + the gather runs: it is the shipped behaviour for the geometries that + declare the capability, gated on measured speed, value, and memory + (2026-08-11, four H100s, both geometries). Setting it to False + selects the banded walk, which stays in place as the rollback. + + The ENVIRONMENT may override the switch either way, which is what + lets one session run both shapes over the same inputs and compare + their values. + """ + if not self.column_gather_geometry: + return False + override = os.environ.get(COLUMN_GATHER_ENV_VAR, '').strip().lower() + if override in _COLUMN_GATHER_ON_VALUES: + return True + if override in _COLUMN_GATHER_OFF_VALUES: + return False + switch = getattr(self, 'forward_column_gather', None) + return True if switch is None else bool(switch) + + def _forward_pixel_batch(self): + """How many pixel columns one gathered cylinder covers. + + :data:`FORWARD_PIXEL_BATCH` carries the value and its provenance. + ``forward_project_pixel_batch`` on the model overrides it, the same + way ``forward_project_slice_band`` overrides the band rule. The + memory ledger calls THIS method rather than re-deriving the number, + so a changed default cannot leave the charge behind.""" + fixed = getattr(self, 'forward_project_pixel_batch', None) + return max(1, int(fixed)) if fixed else FORWARD_PIXEL_BATCH + @staticmethod def _balanced_slice_bounds(extent, band_len): """Tile ``[0, extent)`` into balanced bands no longer than @@ -444,12 +526,19 @@ def _sparse_forward_project_sharded(self, voxel_shards, pixel_indices): Under padding each owner projects only its REAL views (padded views have no angles), and its padded view tail is zero-filled after - assembly, keeping the device form inert end to end.""" + assembly, keeping the device form inert end to end. + + A geometry whose slices spread over a range of detector rows can take + :meth:`_sparse_forward_project_columns` instead, which cuts the + cylinder the other way; :meth:`_column_gather_forward` says when.""" if voxel_shards.placement.is_trivial: return _sharding.Shards( [self.projector_functions._sparse_forward_project_single_device( voxel_shards.tensors[0], pixel_indices)], self.sino_placement) + if self._column_gather_forward(): + return self._sparse_forward_project_columns(voxel_shards, + pixel_indices) sp, rp, view_spans, band_ranges, idx_per = self._banded_setup(pixel_indices) pf = self.projector_functions aligned = self.rows_track_slices @@ -551,6 +640,253 @@ def _sparse_forward_project_sharded(self, voxel_shards, pixel_indices): for t, (_v0, _v1, block) in zip(tensors, view_spans)] return _sharding.Shards(tensors, sp) + def _sparse_forward_project_columns(self, voxel_shards, pixel_indices): + """The multi-device forward as a pixel-batched column gather: each + view-owner walks the pixel axis in batches, gathers each batch's + cylinder at every slice from every slice-owner, and makes ONE + projector call per batch over its own views and the whole slice + range. The alternative to the banded walk in + :meth:`_sparse_forward_project_sharded`, for the geometries + :meth:`_column_gather_forward` admits. + + WHY the shape exists, for the two geometries that take it. A + geometry whose slices spread over a range of detector rows pays per + projector call whatever the slice band contains, because the call's + output spans the whole detector either way. Walking one band per + slice-owner therefore costs that owner count times one full call, and + the forward stops falling when devices are added -- measured flat at + 32.2, 30.6 and 30.5 s over one, two and four devices. A full-height + call per pixel batch is the shape the single-device path already + runs, so the work divides with the view split the way it was meant + to. Measured 2026-08-10 on four H100s, job mg10: cone's per-device + forward fell from 29.7 to 19.4 s at two devices and from 29.3 to + 15.3 s at four, with a lower peak. + + A ROW-ALIGNED geometry's banded walk does divide the work, so its + reason is the other one: the forward kernel is about twice as + efficient per slice on a full-width block of values as on the + shard-width blocks the banded walk hands it, and this shape hands it + full width at every device count. Measured 2026-08-10 on one H100, + at 0.0411 ms per slice on a 1008-wide block against 0.0823 on a + 504-wide one with the device count held at one. + + WHAT DOES NOT MOVE. Every view-owner still produces its own views' + whole sinogram block, from the same voxels, through the same body, so + the operator is unchanged and the sharded forward stays the adjoint + of the sharded back. Only which device assembles which voxels + changes, and the back driver is untouched. Two summation orders do + change for a two-fan geometry: the vertical sum moves from a host-side + sum across bands into the body, and the pixel sum moves the other way, + from the body into the host-side sum across pixel batches. Both sit + inside the value class the forward already has. A row-aligned + geometry has no vertical sum to move, its rows being concatenated + rather than added, so the pixel sum is the whole of what changes + there, and nothing changes at all when one batch covers the pass. + + The two skips of the banded form are kept or dropped deliberately. A + view-owner with no real views receives no gathers and produces an + empty block, as before. The banded form's all-padding sub-band skip + has no counterpart here, because a gathered cylinder spans every + slice-owner at once; the padding it carries is inert (see + :func:`_sharding.gather_column_band`). + + Each batch's gather is issued ONE BATCH AHEAD of the projection that + reads it, and on CUDA its copies run on a stream of their own, so a + device projects one batch while the next batch's values are still + moving to it. Issuing early is what makes that possible; the separate + stream is what makes it happen, because a stream runs its work one + item at a time and copies sharing the projection's stream could only + take turns with it. The batches are summed in the same order they + always were, so the values do not move; what changes is that a device + holds one more cylinder at once, which the memory ledger charges + (COLUMN_GATHER_RESIDENTS). The comment at the gather gives the full + ordering argument, and off CUDA the gather stays the synchronous one it + has always been. + + Each batch after the first adds into the owner's block from INSIDE the + projector's view loop (``accumulate_into``), rather than receiving its + own block for the driver to add. That drops a full-block pass and a + full-block allocation per batch -- a cost that does not shrink with the + batch, so bigger batches only hide it -- and it lowers the widest + instant by the block it no longer allocates. The summation order is + unchanged, element for element. The comment at the accumulation carries + the argument, including why a preallocated zeroed buffer would be worse + rather than better. + + ``forward_project_slice_band`` has nothing to act on here, because + this shape does not band the slice axis at all; what bounds the + transfer instead is the pixel batch. The memory ledger stops + charging the band copy to match. ``back_project_slice_band`` is + unaffected, the back driver being untouched.""" + sp, rp, view_spans, _band_ranges, idx_per = self._banded_setup( + pixel_indices) + pf = self.projector_functions + num_channels = int(self.get_params('sinogram_shape')[2]) + # How tall a block one call returns, which is what the empty blocks + # below have to match. A row-aligned geometry's body sizes its output + # by the values it was handed, and the gathered cylinder is the whole + # DEVICE-form slice axis -- padded tail included, which is exactly the + # length that geometry's sinogram pads its detector rows to. A + # geometry whose slices spread over a range of rows returns the real + # detector rows whatever it is handed. + num_rows = (int(rp.padded_size) if self.rows_track_slices + else int(self.get_params('sinogram_shape')[1])) + num_pixels = int(idx_per[0].shape[0]) + shards = voxel_shards.tensors # in device = global slice order + pixel_batch = self._forward_pixel_batch() + batch_bounds = [(p0, min(p0 + pixel_batch, num_pixels)) + for p0 in range(0, num_pixels, pixel_batch)] + + def worker(i, dev): + v0, v1, _block = view_spans[i] + if v1 <= v0: + # A view-owner with no real views (the sparse-view extension) + # produces an empty block, which assembles as pure zeros. + return torch.zeros((0, num_rows, num_channels), + dtype=voxel_shards.dtype, device=dev) + local_idx = idx_per[i] + owned = None + + def gather(k): + p0, p1 = batch_bounds[k] + return _sharding.gather_column_band_async( + shards, p0, p1, dev, self.dev2dev_safe) + + # The batch after the one being projected, gathered ahead of it. + # A pass of one batch has nothing to gather ahead, and no pixels + # at all leaves this empty. + ahead = gather(0) if batch_bounds else None + for k, (p0, p1) in enumerate(batch_bounds): + full_cyl, ready = ahead + # Issue the NEXT batch's gather before this batch is + # projected, rather than after, so its copies are already + # moving while this projection runs. Nothing here waits for a + # value: run_per_device performs no synchronization, and the + # gather returns once its copies are issued. + # + # THE ORDERING, end to end. Four things arrange it, and each + # covers a different way the copies and the projections could + # get in each other's way. + # + # The copies run on a stream of their own, one per device + # (:func:`_sharding.copy_stream`). A stream runs its work in + # the order it was given, one item at a time, so copies left + # on the stream a device projects on could only take turns + # with the projections, however early they were issued. On + # their own stream the two run at once. + # + # Before any copy starts, each copy stream waits for its + # device's compute stream, so a copy cannot read a shard + # before the kernel that wrote it has finished + # (``open_copy_streams``, called once below). + # + # Every batch carries its OWN event, recorded on the copy + # stream once that batch's copies and their concatenation are + # queued. The compute stream waits for that one event just + # before the projection that reads that batch, so a projection + # never starts on a cylinder that has not arrived -- and never + # waits for the batch gathered ahead of it, which is the work + # meant to be moving right now. + # + # After the pass, each compute stream waits for its copy + # stream (``close_copy_streams``), so a later update cannot + # overwrite a shard while a copy is still reading it. + # + # Off CUDA none of this applies: the gather copies + # synchronously, returns no event, the wait below does + # nothing, and the values are the ones the plain path has + # always produced. + ahead = gather(k + 1) if k + 1 < len(batch_bounds) else None + _sharding.wait_for_column_band(dev, ready) + # THE ACCUMULATION. The first batch's projection allocates the + # owner's block and fills it; every later batch adds into that + # same block from inside the projector's own view loop, which + # is where the block was going to be written anyway. + # + # What this removes, per batch after the first: the projector + # allocated a fresh full block, copied its view batches into + # it, and handed it back for the driver to add -- two full-block + # passes and one full-block allocation where a single pass does + # the same work. The cost is the same at every batch size, so + # it is one the bigger batches HIDE rather than remove, and a + # 1024-class pass at the default batch runs on the order of a + # hundred of them. + # + # NOT a preallocated zeroed buffer, which is the shape this + # looks like from a distance and is strictly worse: adopting + # the first batch's block, as below, costs no zero-fill and no + # add, while a zeroed buffer pays both. + # + # THE VALUES DO NOT MOVE. Per element the sequence is still + # batch 0's contribution, then batch 1's added to it, then batch + # 2's -- the same summands added in the same order as the + # driver-side add did. Only where the addition happens changes, + # so the result is bit for bit what it was. + # + # STREAM LIFETIMES are untouched, and the persistent block needs + # no record_stream. It is allocated, written and added into + # ONLY by this device's compute stream, in program order, and a + # stream runs its work in order; the copy streams read the + # slice-owners' shards and write the gathered cylinders, and + # never touch this block. Holding it across batches instead of + # freeing it each time also keeps its memory out of the caching + # allocator between batches, so it can never be handed to a + # copy stream mid-pass. + if owned is None: + owned = pf.sparse_forward_project_view_range( + full_cyl, local_idx[p0:p1], (v0, v1), slice_start=0, + dev_index=i) + else: + pf.sparse_forward_project_view_range( + full_cyl, local_idx[p0:p1], (v0, v1), slice_start=0, + dev_index=i, accumulate_into=owned) + # Released once the projection that reads it has been issued, so + # a device carries this batch's cylinder no further. With the + # gather ahead of it, the batch after this one is already + # resident by now, which is the third cylinder the memory ledger + # charges (COLUMN_GATHER_RESIDENTS). The release moved after + # the accumulation because the accumulation moved INTO the + # projection; the widest instant is narrower than it was, the + # separate incoming block having gone. + full_cyl = None + if owned is None: + # No pixels at all: the owner still owes its views' block, + # and the banded form would have produced it as zeros too. + owned = torch.zeros((v1 - v0, num_rows, num_channels), + dtype=voxel_shards.dtype, device=dev) + return owned + + # ONE fan-out for the whole call, with the pixel loop inside the + # worker: a fan-out per pixel batch would issue a thread dispatch per + # (batch, device), and putting the loop inside also issues each + # device's gathers from the thread that consumes them. + # + # The copies read the slice-owners' shards and land on the + # view-owners, so both sets of devices have a copy stream to order + # (see the comment at the gather above). Off CUDA both calls do + # nothing. + gather_devices = (list(voxel_shards.placement.devices) + + list(sp.devices)) + _sharding.open_copy_streams(gather_devices) + try: + with self._band_pool(sp.n_devices) as pool: + tensors = _sharding.run_per_device(sp.devices, worker, + executor=pool) + finally: + # Closed even if a worker raised: copies that were already issued + # are still in flight, and the shards they read must not be + # overwritten under them. + _sharding.close_copy_streams(gather_devices) + if sp.is_padded: + # The banded form's own tail fill: zero-fill each owner's padded + # view tail up to its block length. + tensors = [ + t if t.shape[0] == block else torch.cat( + [t, torch.zeros((block - t.shape[0],) + tuple(t.shape[1:]), + dtype=t.dtype, device=t.device)]) + for t, (_v0, _v1, block) in zip(tensors, view_spans)] + return _sharding.Shards(tensors, sp) + def _sparse_back_project_sharded(self, sino_shards, pixel_indices, coeff_power=1): """The banded sharded back (the forward's adjoint): every view-owner @@ -1190,6 +1526,26 @@ def _shard_sinogram(self, sinogram): # silently mis-assemble a geometry that forgot to declare itself. rows_track_slices = False + # Whether this geometry's multi-device forward MAY gather pixel columns + # instead of walking slice bands (see _column_gather_forward). False is + # the base value because the shape has been measured on cone beam and + # parallel beam only: translation and multiaxis have the same + # band-independent per-call cost as cone and should gain from it too, but + # a geometry is switched over on its own measurement rather than on the + # argument. Declaring this True is what lets the gather run, and it runs + # by default -- forward_column_gather = False selects the banded walk. + column_gather_geometry = False + + # The fewest pixels this geometry's COMPILED bodies may be called with. + # 1 -- the base value -- means any width, which is what a geometry whose + # compiled bodies are all correct wants. A geometry that declares more + # gets narrow calls padded up to that width and unpadded again outside the + # compiled region (see projectors.forward_at_min_pixel_width, which also + # carries the measured reason parallel beam declares 2). It is a property + # of the geometry's bodies rather than a user setting, so it is a class + # attribute and not a parameter. + min_compiled_pixel_width = 1 + # Which measured set of widening speed floors governs this geometry's # automatic device count (see _widening_floors). None -- the base value # -- means the parallel floors, which are the more permissive measured diff --git a/mbirtorch/translation_model.py b/mbirtorch/translation_model.py index 29b73d5..24bfb9b 100644 --- a/mbirtorch/translation_model.py +++ b/mbirtorch/translation_model.py @@ -17,8 +17,12 @@ Known scale limit, recorded at port time: at production TCT detector shapes (~1900x3000 panels) the back projection holds (view_batch, P, rows) and (view_batch, P, slices) transients, so large pixel batches are memory-bound -and the view batch shrinks accordingly. A planned engine change may restore -pixel batching; no workaround is built here. +and the view batch shrinks accordingly. What would relieve it is a change +to the projector drivers, not to this file: they currently tile over views +only, and tiling over the pixel axis as well -- the two-axis tiling +described in projectors.py, which mbirjax's sparse projection drivers do -- +would let the pixel batch shrink instead of the view batch. Nothing here +works around its absence. """ import warnings @@ -112,8 +116,9 @@ def _translation_forward_view_batch(values, pixel_indices, view_params_batch, slice_start + L); the z geometry stays anchored on the full num_slices center and taps outside the band contribute zero. - ``plan`` is the memoization slot for a future sorted/CSR stream variant; - unused today.""" + ``plan`` is accepted and ignored. It reserves a place for a future + body that would precompute its geometry once and reuse it across + calls; nothing reads it today.""" n_p, centers, W_p_c, weight_scale, pixel_mag = _translation_horizontal_data( pixel_indices, view_params_batch, num_recon_rows, num_recon_cols, num_channels, delta_voxel, delta_voxel_row, delta_det_channel, @@ -179,8 +184,9 @@ def _translation_back_view_batch(sino_batch, pixel_indices, view_params_batch, gather onto the slices. Returns (P, S), or (P, band_slices) for a slice band, exactly as in cone. - ``plan`` is the memoization slot for a future sorted/CSR stream variant; - unused today.""" + ``plan`` is accepted and ignored. It reserves a place for a future + body that would precompute its geometry once and reuse it across + calls; nothing reads it today.""" n_p, centers, W_p_c, weight_scale, pixel_mag = _translation_horizontal_data( pixel_indices, view_params_batch, num_recon_rows, num_recon_cols, num_channels, delta_voxel, delta_voxel_row, delta_det_channel, diff --git a/mbirtorch/utilities.py b/mbirtorch/utilities.py index 8840e6d..89e8130 100644 --- a/mbirtorch/utilities.py +++ b/mbirtorch/utilities.py @@ -828,17 +828,25 @@ def stitch_arrays(array_list, overlap, axis=2, ramp_overlap=None): return swap(stitched, 0, axis) -def copy_ct_model(ct_model, new_angles=None, new_helical_z_shifts=None, new_num_det_rows=None, new_num_det_cols=None): +def copy_ct_model(ct_model, new_angles=None, new_helical_z_shifts=None, new_num_det_rows=None, new_num_det_cols=None, + new_translation_vectors=None): """ - Create a TomographyModel with the same type and parameters as the given ct_model except with the new input angles - and a corresponding sinogram shape. Restricted to ParallelBeam and ConeBeam models. + Create a TomographyModel with the same type and parameters as the given ct_model except with the new per-view + parameters and a corresponding sinogram shape. Supports the ParallelBeam, ConeBeam, MultiAxisParallel and + Translation models. + + Each geometry names its per-view parameters differently, and the copy uses whichever name the model's own + constructor takes: ``new_angles`` for the three angle-based geometries (a 1D vector for parallel and cone, a + (num_views, 2) array of (azimuth, elevation) pairs for multiaxis) and ``new_translation_vectors`` for + TranslationModel. Passing the argument that does not apply to the given model raises rather than being ignored. If the user explicitly set the devices on ct_model with configure_devices, the copy gets the same devices. Otherwise the copy chooses its own devices when it is used. Args: ct_model (TomographyModel): The model to copy. - new_angles (ndarray of float, optional): 1D vector of projection angles in radians. + new_angles (ndarray of float, optional): Projection angles in radians -- a 1D vector for ParallelBeamModel and + ConeBeamModel, or a (num_views, 2) array of (azimuth, elevation) pairs for MultiAxisParallelModel. If None, then use the angles in ct_model. Defaults to None. new_helical_z_shifts (ndarray of float, optional): 1D vector of per-view axial shifts in ALU for ConeBeamModel. Defaults to None. @@ -846,26 +854,42 @@ def copy_ct_model(ct_model, new_angles=None, new_helical_z_shifts=None, new_num_ If None, then use the num_det_rows in ct_model. Defaults to None. new_num_det_cols (int, optional): Number of detector columns in the new model. If None, then use the num_det_cols in ct_model. Defaults to None. + new_translation_vectors (ndarray of float, optional): (num_views, 3) array of object translations (x, y, z) in + ALU for TranslationModel. If None, then use the translation_vectors in ct_model. Defaults to None. Returns: - An instance of ConeBeamModel or ParallelBeam model + An instance of the same model class as ct_model """ - if str(type(ct_model)).find('ConeBeamModel') > 0: - is_cone = True - elif str(type(ct_model)).find('ParallelBeamModel') > 0: - is_cone = False - else: - raise TypeError('copy_ct_model() supports ConeBeamModel and ParallelBeamModel only; ' - f'got {type(ct_model).__name__}. TranslationModel and ' - 'MultiAxisParallelModel are not yet supported (matching mbirjax); ' - 'construct the new model directly.') + model_name = str(type(ct_model)) + is_cone = model_name.find('ConeBeamModel') > 0 + is_translation = model_name.find('TranslationModel') > 0 + # MultiAxisParallelModel is matched on its own name rather than through 'ParallelBeamModel', which is not a + # substring of it. + if not (is_cone or is_translation or model_name.find('ParallelBeamModel') > 0 + or model_name.find('MultiAxisParallelModel') > 0): + raise TypeError('copy_ct_model() supports ConeBeamModel, ParallelBeamModel, MultiAxisParallelModel and ' + f'TranslationModel; got {type(ct_model).__name__}. Construct the new model directly.') # get_all_params is the single source of truth for reading the params back out: it gives the # constructor args with the view components already unpacked (angles + helical_z_shifts for cone) # and geometry_type in required, so build_model can reconstruct the class. required, optional, regularization = ct_model.get_all_params() - old_angles = required['angles'] + # The key the per-view parameters arrive under is the one the constructor declares, so the copy reads and writes + # that key rather than assuming every geometry has angles. Translation carries translation_vectors and no angles + # at all; the other three carry angles, of one column (parallel, cone) or two (multiaxis). + if is_translation: + view_key, new_view_params = 'translation_vectors', new_translation_vectors + if new_angles is not None or new_helical_z_shifts is not None: + raise ValueError('copy_ct_model: a TranslationModel has per-view translations rather than angles; ' + 'pass new_translation_vectors.') + else: + view_key, new_view_params = 'angles', new_angles + if new_translation_vectors is not None: + raise ValueError('copy_ct_model: new_translation_vectors applies to a TranslationModel only; ' + f'got {type(ct_model).__name__}, so pass new_angles.') + + old_view_params = required[view_key] new_shape = list(required['sinogram_shape']) if is_cone: @@ -884,14 +908,16 @@ def copy_ct_model(ct_model, new_angles=None, new_helical_z_shifts=None, new_num_ raise ValueError('copy_ct_model: new_helical_z_shifts must have the same length as the existing angles.') required['helical_z_shifts'] = new_helical_z_shifts - if new_angles is None: - new_angles = old_angles - new_shape[0] = len(new_angles) + if new_view_params is None: + new_view_params = old_view_params + # len() is the view count for every form here: one entry per view, whether that entry is a scalar angle or a row + # of a (num_views, 2) or (num_views, 3) array. + new_shape[0] = len(new_view_params) if new_num_det_rows is not None: new_shape[1] = new_num_det_rows if new_num_det_cols is not None: new_shape[2] = new_num_det_cols - required['angles'] = new_angles + required[view_key] = new_view_params required['sinogram_shape'] = tuple(new_shape) # The sinogram shape changed, so drop recon_shape and let build_model's auto pass recompute it. @@ -1056,22 +1082,27 @@ def merge_log_files(merged_path, labeled_paths): os.path.abspath(merged_path))) -def get_ct_model(geometry_type, sinogram_shape, angles, source_detector_dist=None, source_iso_dist=None, helical_z_shifts=None): +def get_ct_model(geometry_type, sinogram_shape, angles=None, source_detector_dist=None, source_iso_dist=None, + helical_z_shifts=None, translation_vectors=None): """ Create an instance of TomographyModel with the given parameters Args: - geometry_type (str): 'parallel' or 'cone' + geometry_type (str): 'parallel', 'cone', 'multiaxis' or 'translation' sinogram_shape (tuple list of int): (num_views, num_rows, num_channels) - angles (ndarray of float): 1D vector of projection angles in radians + angles (ndarray of float, optional): Projection angles in radians -- a 1D vector for 'parallel' and 'cone', or a + (num_views, 2) array of (azimuth, elevation) pairs for 'multiaxis'. Not used by 'translation', which takes + translation_vectors instead. Defaults to None. source_detector_dist (float or None, optional): Distance in ALU from source to detector. Defaults to None for geometries that don't need this. source_iso_dist (float or None, optional): Distance in ALU from source to iso. Defaults to None for geometries that don't need this. helical_z_shifts (ndarray, optional): Per-view axial shifts (ALU), same length as angles. Required when use_helical=True. + translation_vectors (ndarray of float, optional): (num_views, 3) array of object translations (x, y, z) in ALU. + Required for geometry_type 'translation' and unused by the others. Defaults to None. Returns: - An instance of ConeBeamModel or ParallelBeam model + An instance of ConeBeamModel, ParallelBeamModel, MultiAxisParallelModel or TranslationModel """ import mbirtorch @@ -1082,11 +1113,21 @@ def get_ct_model(geometry_type, sinogram_shape, angles, source_detector_dist=Non if helical_z_shifts is not None: warnings.warn("Helical mode (helical_z_shifts) is only supported for geometry_type='cone'; ignoring z_shifts.", UserWarning) model = mbirtorch.ParallelBeamModel(sinogram_shape, angles) + elif geometry_type == 'multiaxis': + if helical_z_shifts is not None: + warnings.warn("Helical mode (helical_z_shifts) is only supported for geometry_type='cone'; ignoring z_shifts.", UserWarning) + model = mbirtorch.MultiAxisParallelModel(sinogram_shape, angles) + elif geometry_type == 'translation': + if translation_vectors is None: + raise ValueError("get_ct_model() with geometry_type 'translation' needs translation_vectors, a " + "(num_views, 3) array of object translations in ALU; a translation geometry has no " + "angles.") + model = mbirtorch.TranslationModel(sinogram_shape, translation_vectors, + source_detector_dist=source_detector_dist, + source_iso_dist=source_iso_dist) else: - raise ValueError("get_ct_model() supports geometry_type 'cone' and 'parallel' only; " - f"got {geometry_type!r}. For the translation and multiaxis " - "geometries (not yet supported here, matching mbirjax), construct " - "TranslationModel or MultiAxisParallelModel directly.") + raise ValueError("get_ct_model() supports geometry_type 'cone', 'parallel', 'multiaxis' and 'translation'; " + f"got {geometry_type!r}.") return model @@ -1488,6 +1529,7 @@ class ModelType(str, Enum): PARALLEL = 'parallel' CONE = 'cone' TRANSLATION = 'translation' + MULTIAXIS = 'multiaxis' def generate_demo_data( @@ -1507,6 +1549,7 @@ def generate_demo_data( helical_z_range=None, helical_z_center=0.0, use_curved_detector=False, + elevation_degrees=0.0, voxel_row_aspect=1.0, voxel_slice_aspect=1.0, target_max_attenuation=None, @@ -1544,6 +1587,9 @@ def generate_demo_data( helical_z_range (float, optional): Total axial travel over the scan in ALU for helical mode. helical_z_center (float, optional): Midpoint of axial travel over the scan in ALU for helical mode. use_curved_detector (bool, optional): (cone beam geometry parameter) + elevation_degrees (float, optional): (multiaxis geometry parameter) The + constant tilt of every view out of the horizontal plane, in degrees. + Defaults to 0.0. voxel_row_aspect (float, optional): Aspect ratio for recon rows relative to columns. Defaults to 1.0. voxel_slice_aspect (float, optional): Aspect ratio for recon slices relative to rows. Defaults to 1.0. target_max_attenuation (float, optional): Target max sinogram attenuation for Shepp-Logan phantom. Defaults to None, for which each voxel is in the range [0, 1]. May not be accurate if any detector or voxel dimensions are not 1. @@ -1571,7 +1617,19 @@ def generate_demo_data( # Initialize model - if model_type == ModelType.PARALLEL: + if model_type == ModelType.MULTIAXIS: + # Azimuths over a half rotation, all views at one elevation (tilt). + azimuths = np.linspace(0, np.pi, num_views, endpoint=False) + elevations = np.deg2rad(elevation_degrees) * np.ones(num_views) + angles = np.column_stack([azimuths, elevations]).astype(np.float32) + sinogram_shape = (num_views, num_det_rows, num_det_channels) + ct_model_for_generation = mbirtorch.MultiAxisParallelModel(sinogram_shape, angles) + ct_model_for_generation.set_params(voxel_row_aspect=voxel_row_aspect) + ct_model_for_generation.set_params(voxel_slice_aspect=voxel_slice_aspect) + ct_model_for_generation.auto_set_recon_geometry() + params = {'angles': angles, 'elevation_degrees': elevation_degrees, + 'voxel_row_aspect': voxel_row_aspect, 'voxel_slice_aspect': voxel_slice_aspect} + elif model_type == ModelType.PARALLEL: start_angle = 0 sinogram_shape = (num_views, num_det_rows, num_det_channels) angles = np.linspace(start_angle, end_angle, num_views, endpoint=False) diff --git a/tests/generate_goldens.py b/tests/generate_goldens.py index 094329d..02bd850 100644 --- a/tests/generate_goldens.py +++ b/tests/generate_goldens.py @@ -8,6 +8,17 @@ Writes tests/goldens/golden_.npz (gitignored; regenerate at will). The recorded jax version is the frozen comparison baseline (0.10.1). +WHERE THE GOLDENS COME FROM ON ANOTHER MACHINE. tests/goldens/ is gitignored, +so nothing ships them: a fresh checkout has no archive, and every parity test +in tests/ skips itself with a message naming this script. A run that is meant to +ENFORCE parity -- a nightly, or a release check -- has to run this script in +the mbirjax env FIRST and then run the suite with RUN_GOLDENS=1 (see +dev_scripts/run_tests.sh), because a run that SKIPS these tests reports the +same "passed" as a run that gates them. An archive generated before a +geometry was added is missing that geometry's keys, so its tests skip while +the rest still pass; regenerating after anything is added here is part of the +same step. + Contents per cell: the shepp-logan phantom, its sinogram, transmission-root weights, sparse fwd/back outputs on a fixed subset, the qGGMRF gradient and Hessian on that subset, the Hessian diagonal, the FBP recon, the auto-set diff --git a/tests/test_demo_data.py b/tests/test_demo_data.py index e3b98e1..e506d6f 100644 --- a/tests/test_demo_data.py +++ b/tests/test_demo_data.py @@ -155,3 +155,13 @@ def test_gen_translation_vectors_grid(): assert np.allclose(vecs[:, 1], 0.0) # no y motion assert np.allclose(sorted(set(vecs[:, 0])), [-10.0, 0.0, 10.0]) assert np.allclose(sorted(set(vecs[:, 2])), [-2.5, 2.5]) + + +def test_generate_demo_data_multiaxis(): + phantom, sino, params = mbirtorch.generate_demo_data( + model_type='multiaxis', elevation_degrees=25.0, object_type='cube', + num_views=8, num_det_rows=16, num_det_channels=24) + assert sino.shape == (8, 16, 24) + assert params['angles'].shape == (8, 2) + assert np.allclose(params['angles'][:, 1], np.deg2rad(25.0)) + assert np.isfinite(np.asarray(sino)).all() and np.asarray(sino).max() > 0 diff --git a/tests/test_denoiser.py b/tests/test_denoiser.py index 88fe627..246ae18 100644 --- a/tests/test_denoiser.py +++ b/tests/test_denoiser.py @@ -71,8 +71,8 @@ def test_denoise_reduces_noise(device): def test_sharded_denoise_matches_single_device(): """Two CPU shards vs one device on the same seeded problem. The sharded path stages halos once per pass and combines the step-size sums on the - host, so agreement is at float level, not bitwise (gate per the measured - iterated-comparison floor).""" + lead device, so agreement is at float level, not bitwise (gate per the + measured iterated-comparison floor).""" shape = (24, 24, 21) # 2 shards pad the slice axis 21 -> 22 clean = np.zeros(shape, dtype=np.float32) clean[6:-6, 6:-6, 5:-5] = 1.0 diff --git a/tests/test_device_policy.py b/tests/test_device_policy.py index 4fc84f6..aa73a10 100644 --- a/tests/test_device_policy.py +++ b/tests/test_device_policy.py @@ -37,6 +37,25 @@ def make_model(shape=(8, 6, 8), device='cpu', **kwargs): return model +@pytest.fixture(autouse=True) +def kernel_declared_projection(monkeypatch): + """Price the projection the way the CUDA model these tests stand in for + would price it. + + A CUDA parallel or cone model binds the hand-written kernel bodies, and + each of those declares what one of its views holds. On CPU no kernel is + available, so the same model binds the general torch bodies instead, and + the ledger prices a torch body's views for itself at a much larger + residency (``_memory_ledger.TORCH_BODY_VIEW_SLABS``). These tests are + about the device-count RULE, not about either residency, so they hold the + projection charge at the kernel-declared one; otherwise the capacity + arithmetic they drive would be a different model's. The torch-body + charge has its own tests in test_memory_ledger.py. + """ + monkeypatch.setattr(_memory_ledger, 'torch_body_directions', + lambda model: ()) + + @pytest.fixture def no_speed_guard(monkeypatch): """Turn off the widening speed floors. @@ -698,44 +717,92 @@ def test_skipping_the_memory_preflight_leaves_the_speed_floors_in_force( # ── geometries the floors have never met ───────────────────────────────────── class _UnlistedGeometry(mbirtorch.ParallelBeamModel): - """A geometry that never declared a floor family, as TranslationModel and - any future class would arrive.""" + """A stand-in that declares no floor family. + + The two real classes below are the standing coverage for this path, and + this one is kept beside them for the same reason + test_widening_floors.py keeps a synthetic table: the RULE has to outlive + the data. Once a refresh measures multiaxis and translation they will + declare families of their own and stop exercising the fallback, and the + fallback still has to work for whatever geometry arrives next. + """ _floor_family = None -def test_a_model_with_no_floor_family_gets_the_parallel_floors(monkeypatch, - unpinned): - def make(shape): - angles = np.linspace(0, np.pi, shape[0], endpoint=False) - model = _UnlistedGeometry(shape, angles) - model.configure_devices(devices=['cpu']) - model.set_params(no_warning=True, verbose=0) - return model - - small = with_four_visible(monkeypatch, make(CELL_128)) +def _synthetic_no_family(shape): + angles = np.linspace(0, np.pi, shape[0], endpoint=False) + return _UnlistedGeometry(shape, angles) + + +def _multiaxis_no_family(shape): + """Multiaxis angles are (azimuth, elevation) pairs, one row per view.""" + azimuth = np.linspace(0, np.pi, shape[0], endpoint=False) + elevation = np.linspace(-0.4, 0.4, shape[0]) + return mbirtorch.MultiAxisParallelModel(shape, np.stack([azimuth, elevation], axis=1)) + + +def _translation_no_family(shape): + """Translation views are object translations, laid out on a grid whose + two side lengths multiply to the view count.""" + num_views = shape[0] + num_x = 16 if num_views == CELL_128[0] else 32 + vectors = mbirtorch.gen_translation_vectors(num_x, num_views // num_x, + x_spacing=3.0, z_spacing=2.0) + return mbirtorch.TranslationModel(shape, vectors, + source_detector_dist=4.0 * shape[2], + source_iso_dist=1.0 * shape[2]) + + +UNMEASURED_GEOMETRIES = [ + (_multiaxis_no_family, 'MultiAxisParallelModel'), + (_translation_no_family, 'TranslationModel'), + (_synthetic_no_family, '_UnlistedGeometry'), +] + + +def _built(make, shape, verbose=0): + model = make(shape) + model.configure_devices(devices=['cpu']) + model.set_params(no_warning=True, verbose=verbose) + return model + + +@pytest.mark.parametrize("make,class_name", UNMEASURED_GEOMETRIES, + ids=[name for _make, name in UNMEASURED_GEOMETRIES]) +def test_a_model_with_no_floor_family_gets_the_parallel_floors( + monkeypatch, unpinned, make, class_name): + """Every class that declares no floor family is governed by the parallel + floors -- checked on the two real geometries that arrive that way, not + only on a stand-in.""" + model = make(CELL_128) + assert type(model).__name__ == class_name + assert model._floor_family is None + + small = with_four_visible(monkeypatch, _built(make, CELL_128)) small._apply_device_policy() assert small.sino_placement.n_devices == 1 # The permissive set, not a refusal: at the parallel n=2 floor it widens. - at_the_floor = with_four_visible(monkeypatch, make(CELL_512)) + at_the_floor = with_four_visible(monkeypatch, _built(make, CELL_512)) at_the_floor._apply_device_policy() assert at_the_floor.sino_placement.n_devices == 2 +@pytest.mark.parametrize("make,class_name", UNMEASURED_GEOMETRIES, + ids=[name for _make, name in UNMEASURED_GEOMETRIES]) def test_the_substituted_family_is_named_in_the_log(monkeypatch, unpinned, - caplog): + caplog, make, class_name): """A geometry that was never measured must not have that fact hidden from - it, so the selection path says which floors it borrowed.""" - angles = np.linspace(0, np.pi, CELL_128[0], endpoint=False) - model = _UnlistedGeometry(CELL_128, angles) - model.configure_devices(devices=['cpu']) - model.set_params(no_warning=True, verbose=2) - with_four_visible(monkeypatch, model) + it, so the selection path says which class borrowed which floors.""" + model = with_four_visible(monkeypatch, _built(make, CELL_128, verbose=2)) with caplog.at_level('DEBUG', logger=model.logger.name): model._apply_device_policy() assert 'names no _floor_family' in caplog.text assert 'parallel widening speed floors' in caplog.text + # The line names the class, so a log read months later says which + # geometry was running on borrowed numbers. + assert class_name in caplog.text # ── the split_sino_recon halves ────────────────────────────────────────────── diff --git a/tests/test_memory_ledger.py b/tests/test_memory_ledger.py index c7ea12d..61a7ca5 100644 --- a/tests/test_memory_ledger.py +++ b/tests/test_memory_ledger.py @@ -141,29 +141,85 @@ def test_persistent_set_shrinks_with_the_device_count(): assert persistent_4 == persistent_1 // 4 -def test_band_reduce_is_flat_in_the_device_count(): - """The finding the error message's remedy ordering rests on. - - sum_band_to_owner materializes all n partials on the owner before summing, - and one band is the whole shard by default, so the term reads about 1.5x a - full-volume cylinder set at both two and four devices. Adding devices - shrinks the persistent set and leaves this where it was. +def test_band_reduce_shrinks_with_the_device_count(): + """The signature that replaced the flat one, and the closed form it rests + on. + + sum_band_to_owner used to move all n partials onto the owner before + summing them, so the owner held n bands plus the running total. One band + is the whole shard by default, so that was about 1.5x a full-volume + cylinder set at BOTH two and four devices: adding devices shrank the + persistent set and left this where it was. The reduce now streams each + arriving partial in bounded row slabs, so the owner holds its running + total, the partial it produced itself, and one slab per source -- + ``num_pixels x (shard + band)`` plus ``(n - 1)`` slabs, which is two + cylinder-SHARDS at the default band and therefore halves when the device + count doubles. + + Priced at a production-like size, where a band is far larger than one + slab. At the small sizes the other tests use, a whole band fits inside a + single slab and moves in one piece, exactly as it always did. """ + pixels, slices = 800_000, 1024 + def reduce_bytes(n): - ledger = estimate_peak_device_bytes(make_plan(n_devices=n)) + ledger = estimate_peak_device_bytes(make_plan( + n_devices=n, recon=(1024, 1024, slices), + num_pixels_full=pixels, granularities=(1,))) return dict(_sub(ledger, 'subset back projection', n, 'band reduce').terms)['band reduce'][0] + # The closed form, pinned exactly at both counts. + for n in (2, 4): + band = slices // n # one shard, the default band + slab = _sharding.reduce_slab_rows(pixels, band * 4) * band * 4 + assert reduce_bytes(n) == 2 * pixels * band * 4 + (n - 1) * slab two, four = reduce_bytes(2), reduce_bytes(4) - assert two > 0 and four > 0 - # Flat within 1 percent between n=2 and n=4, not shrinking like 1/n. - assert abs(two - four) / two < 0.01 - # And it does NOT collapse toward zero: both sit near 1.5x one full set. - subset_full = math.ceil(800 / 4) * 32 * 4 - assert two == pytest.approx(1.5 * subset_full, rel=0.02) + # It now falls with the device count instead of standing still. Not + # exactly a half, because the slab term is a fixed number of bytes and + # there is one more of them at four devices. + assert 0.5 <= four / two <= 0.56 + # And it is well under what the old materialize-then-sum form charged: + # n + 1 bands at two devices, n + 2 at four. + assert two < 0.8 * 3 * pixels * (slices // 2) * 4 + assert four < 0.4 * 6 * pixels * (slices // 4) * 4 assert reduce_bytes(1) == 0 # a single device never runs the reduce +def test_band_reduce_charges_the_bands_already_reduced_this_pass(): + """A band smaller than the shard means several reduces per owner, and the + owner holds the ones it has finished until it concatenates them. + + The old charge counted only the band in flight, so it fell toward zero as + the band narrowed while the owner really was holding most of a shard. + The ``shard + band`` form covers both: the bands already done, at most + ``shard - band``, and the two live ones. + """ + plan = make_plan(n_devices=2, back_band=4) + ledger = estimate_peak_device_bytes(plan) + charged = dict(_sub(ledger, 'subset back projection', 2, + 'band reduce').terms)['band reduce'][0] + p_sub, shard, band = math.ceil(800 / 4), 16, 4 + slab = _sharding.reduce_slab_rows(p_sub, band * 4) * band * 4 + assert charged == p_sub * (shard + band) * 4 + slab + + # The floor rule in the place it bites: however narrow the band, the owner + # still ends the pass holding a whole shard, so the charge may not fall + # under one cylinder-shard. The old form did, which is what this + # replaces: it charged only the band in flight. + def charge(band_length): + led = estimate_peak_device_bytes(make_plan(n_devices=2, + back_band=band_length)) + return dict(_sub(led, 'subset back projection', 2, + 'band reduce').terms)['band reduce'][0] + + one_shard = p_sub * shard * 4 + for band_length in (1, 2, 4, 8, 16): + assert charge(band_length) > one_shard, band_length + # And it still falls as the band narrows, so the knob remains a lever. + assert charge(1) < charge(4) < charge(16) + + def test_empty_shard_extensions_skip_their_role_terms(): """A device with no real views does no projection; one with no real slices holds no band. The ledger charges each role only where it exists.""" @@ -851,6 +907,414 @@ def test_recon_is_unaffected_on_a_cpu_model(): assert np.all(np.isfinite(recon)) +# ── the torch-body projection charge ───────────────────────────────────────── +# A torch body is a projection body written as general torch code, which is +# what a geometry with no hand-written kernel runs. It declares no per-view +# cost, so the ledger prices its views itself. +SLABS = _memory_ledger.TORCH_BODY_VIEW_SLABS + + +def test_torch_body_directions_follow_the_bound_bodies(): + """A body that declares its own per-view cost is priced by that + declaration; one that declares nothing is a torch body. The two + directions are asked separately, because a model may bind a kernel one + way and a torch body the other.""" + def kernel_body(): + pass + kernel_body._view_batch_cost = lambda p, cols, args: (1, 1) + + def torch_body(): + pass + + class FakeModel: + def __init__(self, fwd, back): + self._bodies = (fwd, back) + + def _view_batch_bodies(self): + return self._bodies + + directions = _memory_ledger.torch_body_directions + assert directions(FakeModel(kernel_body, kernel_body)) == () + assert directions(FakeModel(torch_body, torch_body)) == ('forward', 'back') + assert directions(FakeModel(torch_body, kernel_body)) == ('forward',) + assert directions(FakeModel(kernel_body, torch_body)) == ('back',) + + +def test_a_declared_per_view_cost_is_charged_exactly_as_declared(): + """The kernel-declared path may not move: a body that states what one of + its views holds is charged that and nothing more.""" + def charge(direction, num_pixels, band_cols): + return 8, 1024 # 8 views at 1 KiB each + + ledger = estimate_peak_device_bytes(make_plan(view_charge=charge)) + terms = dict(_named(ledger, 'back projection').terms) + assert terms['back batch'][0] == 8 * 1024 + forward = dict(_named(ledger, 'initial forward projection').terms) + assert forward['forward batch'][0] == 8 * 1024 + + +def test_a_torch_body_view_batch_is_charged_at_the_measured_slab_count(): + """A torch body holds a loop of slabs where the driver's nominal charge + prices one, so the ledger charges the measured count of them. + + The slab is (view batch, pixels, width) floats, with width the wider of + the detector rows and the slice band the call was handed -- the two axes + the body sweeps. The view batch stays the driver's own choice. + """ + rows, channels, slices = 32, 32, 32 + p_sub = math.ceil(800 / 4) + + def charge(direction, num_pixels, band_cols): + return 8, 1024 # the driver's batch and nominal + + plan_kwargs = dict(view_charge=charge, num_pixels_full=800, + num_rows=rows, num_channels=channels, + recon=(32, 32, slices)) + declared = estimate_peak_device_bytes(make_plan(**plan_kwargs)) + torch_body = estimate_peak_device_bytes(make_plan( + torch_body_directions=('forward', 'back'), **plan_kwargs)) + + # One device: the band is the whole slice axis, so width is max(32, 32). + width = max(rows, slices) + back = dict(_named(torch_body, 'back projection').terms)['back batch'][0] + assert back == SLABS * 8 * p_sub * width * 4 + assert dict(_named(declared, 'back projection').terms)['back batch'][0] \ + == 8 * 1024 + forward = dict(_named(torch_body, 'initial forward projection') + .terms)['forward batch'][0] + assert forward == SLABS * 8 * 800 * width * 4 + # Every other term is untouched, so the peak moves only by the charge. + assert torch_body.peak_bytes(0) > declared.peak_bytes(0) + + +def test_the_torch_body_slab_follows_the_wider_of_rows_and_band(): + """The body allocates arrays at the detector-row extent AND at the slice + band; the wider of the two sets the slab. Under sharding the band is one + owner's shard, so a tall volume's slab shrinks with the device count and a + wide detector's does not.""" + def charge(direction, num_pixels, band_cols): + return 1, 1 + + def batch(rows, slices, n_devices): + ledger = estimate_peak_device_bytes(make_plan( + n_devices=n_devices, view_charge=charge, + torch_body_directions=('forward', 'back'), + num_pixels_full=800, num_rows=rows, recon=(32, 32, slices))) + return dict(_named(ledger, 'initial forward projection') + .terms)['forward batch'][0] + + # Tall volume, narrow detector: at one device the band is all 64 slices, + # and at four devices it is the 16-slice shard -- below the 32 rows, which + # then set the slab. + assert batch(32, 64, 1) == SLABS * 800 * 64 * 4 + assert batch(32, 64, 4) == SLABS * 800 * 32 * 4 + # Wide detector: the rows set the slab at every device count. + assert batch(128, 32, 1) == SLABS * 800 * 128 * 4 + assert batch(128, 32, 4) == SLABS * 800 * 128 * 4 + + +def test_a_torch_body_pays_for_both_forward_blocks(): + """The forward loop holds the outgoing block and the incoming one. + + Against a body that declares its own cost, one of the two is already + inside the batch charge, because a forward kernel body's declaration + prices its output plane per view. A torch body declares nothing, and + what the ledger charges in its place is the body's internal slab set, + which does not include the output plane -- so both blocks are charged. + """ + rows, channels = 32, 32 + + def charge(direction, num_pixels, band_cols): + return 8, 1024 # 4 batches over 32 views + + def block(directions): + ledger = estimate_peak_device_bytes(make_plan( + n_devices=2, view_charge=charge, + torch_body_directions=directions)) + return dict(_named(ledger, 'initial forward projection') + .terms)['forward block'][0] + + assert block(()) == 1 * 8 * rows * channels * 4 + assert block(('forward', 'back')) == 2 * 8 * rows * channels * 4 + # The back direction alone leaves the forward's own term where it was. + assert block(('back',)) == 1 * 8 * rows * channels * 4 + + +# One row per measured arm: (sinogram shape, recon shape, masked pixel count, +# per-device measured peak bytes). Measured 2026-08-10 on four H100s (job +# mg8) -- the two geometries with no hand-written kernels, at one, two and +# four devices, weighted, from a supplied sinogram with no initial volume. +MEASURED_ARMS = { + 'ma1024_n1': ((1024, 1008, 992), (992, 992, 1148), 771240, + [37310451712]), + 'ma1024_n2': ((1024, 1008, 992), (992, 992, 1148), 771240, + [26138702848, 23767433216]), + 'ma1024_n4': ((1024, 1008, 992), (992, 992, 1148), 771240, + [17888278016, 17820163072, 17820163072, 16934779392]), + 'ma512_n1': ((512, 448, 384), (384, 384, 510), 115164, + [12253271552]), + 'ma512_n2': ((512, 448, 384), (384, 384, 510), 115164, + [9492893184, 9452805632]), + 'ma512_n4': ((512, 448, 384), (384, 384, 510), 115164, + [3768753664, 3757754368, 3757754368, 3757981696]), + 'tct2k_n1': ((256, 1900, 3000), (118, 360, 240), 42480, + [29262431744]), + 'tct2k_n2': ((256, 1900, 3000), (118, 360, 240), 42480, + [40081962496, 17882421760]), + 'tct2k_n4': ((256, 1900, 3000), (118, 360, 240), 42480, + [34051462656, 34081272320, 34081272320, 34081102336]), + 'tct1k_n1': ((256, 950, 1500), (59, 180, 120), 10620, + [8227791872]), + 'tct1k_n2': ((256, 950, 1500), (59, 180, 120), 10620, + [11882288128, 5844949504]), + 'tct1k_n4': ((256, 950, 1500), (59, 180, 120), 10620, + [29162906112, 29161971200, 29161972224, 29161972224]), +} +# The granularity list those runs used, which is the library default. +MEASURED_GRANULARITY = (1, 2, 4, 8, 16, 32, 64, 128, 128, 128, 128) +MEASURED_VISITED = (4, 16, 64) + + +def _measured_view_charge(sinogram_shape, recon_shape, n_devices): + """The view batch and nominal slab the DRIVER chose in those runs. + + Written out here rather than taken from a live model because the batch + depends on the transient budget, and that budget is scaled by the + per-device sinogram on CUDA and flat on CPU -- these tests run on CPU, so + a CPU model would choose a different batch than the measured runs did and + the comparison would be against the wrong arithmetic. + """ + from mbirtorch.projectors import Projectors + views, rows, channels = sinogram_shape + cols = max(int(recon_shape[2]), int(rows)) + local_views = -(-int(views) // int(n_devices)) + budget = max(Projectors.VIEW_BATCH_TRANSIENT_FLOOR_BYTES, + min(Projectors.VIEW_BATCH_TRANSIENT_BUDGET_BYTES, + Projectors.VIEW_BATCH_SINO_MULTIPLE + * local_views * rows * channels * 4)) + + def charge(direction, num_pixels, band_cols): + bytes_per_view = int(num_pixels) * cols * 4 + return (max(1, min(Projectors.VIEW_BATCH_BODY_DEFAULT, + budget // max(1, bytes_per_view))), + bytes_per_view) + return charge + + +def _measured_arm_ledger(arm): + sinogram_shape, recon_shape, num_pixels, measured = MEASURED_ARMS[arm] + n_devices = len(measured) + devices = ['cpu'] * n_devices + sino = _sharding.Placement(devices, axis=0, real_size=sinogram_shape[0]) + recon = _sharding.Placement(devices, axis=-1, real_size=recon_shape[2]) + plan = LedgerPlan( + sinogram_shape=sinogram_shape, + recon_shape=recon_shape, + devices=devices, + view_blocks=[(e - s, v) for _d, (s, e), v + in sino.padded_shard_ranges()], + slice_blocks=[(e - s, v) for _d, (s, e), v + in recon.padded_shard_ranges()], + sino_rows=sinogram_shape[1], + rows_track_slices=False, + num_pixels_full=num_pixels, + num_pixels_grid=recon_shape[0] * recon_shape[1], + granularities=MEASURED_VISITED, + partition_granularities=MEASURED_GRANULARITY, + weights_supplied=True, + # The translation arms carry no cylindrical mask, so their masked set + # IS the whole grid and their hessian back-projects the grid directly. + hessian_masked=num_pixels < recon_shape[0] * recon_shape[1], + view_charge=_measured_view_charge(sinogram_shape, recon_shape, + n_devices), + torch_body_directions=('forward', 'back')) + return estimate_peak_device_bytes(plan), measured + + +@pytest.mark.parametrize('arm', sorted(MEASURED_ARMS)) +def test_the_torch_body_ledger_covers_every_measured_peak(arm): + """The floor, on the runs the slab count was calibrated from. + + A modeled peak below the measured one lets a doomed reconstruction start + and die inside the allocator, which is the failure this module exists to + prevent. Every device of every measured arm must sit at or above 1.00. + """ + ledger, measured = _measured_arm_ledger(arm) + for i, peak in enumerate(measured): + assert ledger.peak_bytes(i) >= peak, ( + f'{arm} device {i}: modeled {ledger.peak_bytes(i)} < ' + f'measured {peak}') + + +def test_the_torch_body_over_charge_stays_inside_its_band(): + """The other side of the floor: an over-charge spreads a reconstruction + over more devices than it needs, so the band is asserted too. It is + wider than CALIBRATION_BAND because one slab count covers two geometries + that hold different numbers of slabs, and because two measured two-device + runs peaked twice as high on one device as on the other from identical + shards. + """ + low, high = _memory_ledger.TORCH_BODY_CALIBRATION_BAND + assert low == _memory_ledger.CALIBRATION_BAND[0] + worst = 0.0 + for arm in MEASURED_ARMS: + ledger, measured = _measured_arm_ledger(arm) + for i, peak in enumerate(measured): + worst = max(worst, ledger.peak_bytes(i) / peak) + assert low <= worst <= high + + +def test_format_calibration_judges_against_the_band_it_is_given(): + rows = [('cuda:0', 40 * GB, 10 * GB, 4.00)] + assert 'over' in _memory_ledger.format_calibration(rows) + assert 'over' not in _memory_ledger.format_calibration( + rows, band=_memory_ledger.TORCH_BODY_CALIBRATION_BAND) + + +# ── the forward's column gather ────────────────────────────────────────────── +@pytest.mark.parametrize('aligned', (False, True), + ids=('two-fan', 'row-aligned')) +def test_the_column_gather_swaps_the_band_copy_for_a_gathered_cylinder(aligned): + """The two states of the same phase. Walking slice bands leaves a + broadcast band resident on every view-owner; gathering columns leaves a + cylinder that is one pixel batch wide and the whole slice axis tall, and + no band at all. Both forward phases carry the swap. + + Both GEOMETRIES are priced by the same arithmetic, and the parametrization + is the claim: what a gather holds is set by the shape it assembles -- one + pixel batch by the whole device-form slice axis -- and not by whether the + geometry's detector rows track its slices. The two take the path for + different reasons and pay the same term for it. + + THREE such cylinders are charged, not one: the driver gathers one batch + ahead of the projection that reads it, so the widest instant holds the + cylinder about to be projected, the pieces arriving for the batch after + it, and their concatenation. The count is written out here rather than + read from the module, so that changing the constant alone cannot move the + charge without this test noticing.""" + slices, batch = 32, 100 # make_plan's slice axis + banded = estimate_peak_device_bytes( + make_plan(n_devices=2, rows_track_slices=aligned)) + gathered = estimate_peak_device_bytes( + make_plan(n_devices=2, rows_track_slices=aligned, + column_pixel_batch=batch)) + for fragment in ('initial forward projection', + 'subset delta forward projection'): + walked = dict(_named(banded, fragment).terms) + columns = dict(_named(gathered, fragment).terms) + assert walked['broadcast band'][0] > 0, fragment + assert walked['column cylinder'] == [0, 0], fragment + assert columns['broadcast band'] == [0, 0], fragment + assert columns['column cylinder'] == [3 * batch * slices * 4] * 2, \ + fragment + + +def test_the_gathered_cylinder_is_capped_by_the_pass_it_covers(): + """A batch wider than the pixel set gathers the pixel set: the charge + follows what one call is actually handed, which is what keeps the term + honest at the small end without a separate rule. + + Such a pass runs as a single batch and so gathers nothing ahead, holding + two cylinders where the charge is three. That over-charge is deliberate: + the ledger's one hard rule is that it may never charge less than a run + needs, and one term that covers the widest instant is simpler than a + second rule for the passes that fall short of it.""" + slices, pixels = 32, 800 + ledger = estimate_peak_device_bytes( + make_plan(n_devices=2, column_pixel_batch=10 ** 6)) + terms = dict(_named(ledger, 'initial forward projection').terms) + assert terms['column cylinder'] == [3 * pixels * slices * 4] * 2 + + +def test_the_gathered_cylinder_does_not_grow_with_the_device_count(): + """The property that dissolves the objection to assembling whole + cylinders: the term is the batch by the WHOLE slice axis on every + view-owner, so adding devices does not change it -- where the broadcast + band it replaces is a shard and halves with the count.""" + charges, bands = [], [] + for n in (2, 4): + gathered = estimate_peak_device_bytes( + make_plan(n_devices=n, column_pixel_batch=100)) + walked = estimate_peak_device_bytes(make_plan(n_devices=n)) + charges.append(dict(_named(gathered, 'initial forward projection') + .terms)['column cylinder'][0]) + bands.append(dict(_named(walked, 'initial forward projection') + .terms)['broadcast band'][0]) + assert charges[0] == charges[1] + assert bands[1] == bands[0] // 2 + # A single device never gathers: it holds the whole volume already. + one = estimate_peak_device_bytes( + make_plan(n_devices=1, column_pixel_batch=100)) + assert dict(_named(one, 'initial forward projection') + .terms)['column cylinder'] == [0] + + +def test_the_column_gather_prices_the_call_it_actually_makes(): + """The two terms that move with the new call shape. One call is handed + the WHOLE device-form slice axis instead of a band, and one pixel batch + instead of every pixel of the pass, so the per-view cost model must be + asked those two numbers.""" + asked = [] + + def charge(direction, num_pixels, band_cols): + asked.append((direction, num_pixels, band_cols)) + return 4, 1024 + + batch = 100 + estimate_peak_device_bytes(make_plan(n_devices=2, view_charge=charge)) + walked = [(p, c) for d, p, c in asked if d == 'forward'] + asked.clear() + estimate_peak_device_bytes( + make_plan(n_devices=2, view_charge=charge, column_pixel_batch=batch)) + gathered = [(p, c) for d, p, c in asked if d == 'forward'] + assert {c for _p, c in walked} == {16} # one slice shard of 32 + assert {c for _p, c in gathered} == {32} # the whole slice axis + assert max(p for p, _c in walked) == 800 # the whole pass + assert max(p for p, _c in gathered) == batch + + +def test_plan_from_model_reads_the_resolved_pixel_batch(monkeypatch): + """The ledger must not re-derive the driver's rule. It asks the model + for the batch it would actually walk, so a changed default or an override + reaches the charge without a second edit here. + + The environment knob is cleared first, because the first assertion reads + the default and a suite run may be forcing the path on around it.""" + from mbirtorch.tomography_model import (COLUMN_GATHER_ENV_VAR, + FORWARD_PIXEL_BATCH) + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + cell = (8, 8, 8) + angles = np.linspace(0, 2 * np.pi, cell[0], endpoint=False) + model = mbirtorch.ConeBeamModel(cell, angles, source_detector_dist=32, + source_iso_dist=16) + model.configure_devices(devices=['cpu']) + model.set_params(no_warning=True, verbose=0) + devices = ['cpu', 'cpu'] + # Unset means the gather (the shipped default), so the charge is present + # at the shipped batch; refusing the gather removes it. + assert _memory_ledger.plan_from_model( + model, devices).column_pixel_batch == FORWARD_PIXEL_BATCH + model.forward_column_gather = False + assert _memory_ledger.plan_from_model( + model, devices).column_pixel_batch is None + model.forward_column_gather = True + model.forward_project_pixel_batch = 512 + assert _memory_ledger.plan_from_model( + model, devices).column_pixel_batch == 512 + # The row-aligned geometry takes the same path, so the same resolution has + # to reach its charge -- present by default, absent when refused, exactly + # as on cone. + par = mbirtorch.ParallelBeamModel(cell, np.linspace(0, np.pi, cell[0], + endpoint=False)) + par.configure_devices(devices=['cpu']) + par.set_params(no_warning=True, verbose=0) + assert _memory_ledger.plan_from_model( + par, devices).column_pixel_batch == FORWARD_PIXEL_BATCH + par.forward_column_gather = False + assert _memory_ledger.plan_from_model( + par, devices).column_pixel_batch is None + + # ── helpers ────────────────────────────────────────────────────────────────── def _named(ledger, fragment): for phase in ledger.phases: diff --git a/tests/test_multiaxis.py b/tests/test_multiaxis.py index 2ae1639..8526209 100644 --- a/tests/test_multiaxis.py +++ b/tests/test_multiaxis.py @@ -2,12 +2,25 @@ goldens against mbirjax (single ops, FBP, auto geometry, and seeded convergence parity), a recon smoke, and 2-shard vs 1-device parity. -Iterated-comparison gates for this geometry are set from its MEASURED parity -floor, not copied from other geometries: at the dividing case (16 views, -elevations to 29 deg) the seeded 3-iteration recon differs from mbirjax by -1.2e-3 max, decaying to 4.2e-4 by 10 iterations -- trajectory float noise -around one fixed point, the same recorded pattern as parallel 1024. Traces -(fm_rmse, alpha) match at 3e-6 / 2e-5. Gates: traces tight, volumes 5e-3. +The two seeded-reconstruction gates are each set from the parity MEASURED at +the configuration that gate runs on, rather than sharing one number, because +the two configurations differ by more than an order of magnitude: + + * The GOLDEN configuration (24 views, elevations to +-0.4 rad) matches + mbirjax to 1.1e-5 max on the volume at 3 iterations, decaying to 6.8e-6 + by 10. Its volume gate is 2e-4, about 18x the measured value. + * The SHARDED comparison runs the dividing case (16 views, elevations to + 29 deg), where three VCD iterations amplify float summation-order + differences of order 1e-7 into 9.4e-4 between 2 shards and 1 device -- + trajectory float noise around one fixed point, the same recorded pattern + as parallel 1024, and the same size as this configuration's own 1.2e-3 + difference from mbirjax at 3 iterations (4.2e-4 by 10). Its volume gate + stays 5e-3, a 5.3x margin over that measurement. + +The golden test's per-iteration traces (alpha, fm_rmse) measure about 6.5e-6 +and 4.7e-6 and are gated further above that than the volume is: a trace is one +scalar per iteration, so a single late step size can move without the +reconstruction moving with it. """ import glob @@ -90,8 +103,15 @@ def test_multiaxis_recon_smoke(device): def test_multiaxis_sharded_recon_matches_single_device(): - """2 CPU shards vs 1 device on the same seeded problem, gated at this - geometry's measured parity floor (see the module docstring).""" + """2 CPU shards vs 1 device on the same seeded problem. + + This runs the dividing configuration, where the reconstruction + trajectory amplifies float summation-order differences: the measured + spread is 9.4e-4, so the 5e-3 gate below is a 5.3x margin. That is a + much looser number than the golden test's, and deliberately so -- see + the module docstring for why the two configurations cannot share one + tolerance. + """ ref_m = _small_ma(['cpu']) rs = ref_m.get_params('recon_shape') phantom = mbirtorch.gen_translation_phantom(rs, 'dots', None, fill_rate=0.05) @@ -174,6 +194,15 @@ def test_multiaxis_fbp(golden, ma_model): @pytest.mark.goldens @ma_golden def test_multiaxis_recon_convergence_parity(golden, ma_model): + """Seeded 3-iteration parity with mbirjax on the GOLDEN configuration. + + The volume gate is set from what this configuration measures, not from + the sharded test's number: 24 views with elevations to +-0.4 rad agree + with mbirjax to 1.1e-5 at 3 iterations and 6.8e-6 at 10, so 2e-4 is + about 18x the measurement -- room for another platform's arithmetic, + while still catching a regression an order of magnitude smaller than the + 5e-3 this test used to share with the sharded comparison. + """ np.random.seed(int(golden["recon_seed"])) recon, rd = ma_model.recon(golden["ma_sino"], max_iterations=3, stop_threshold_change_pct=0.0, logfile_path=None) @@ -187,4 +216,4 @@ def test_multiaxis_recon_convergence_parity(golden, ma_model): f"fm rel = {fm_rel:.2e}, final rel_max = {final_rel:.2e}") assert alpha_rel < 1e-2 assert fm_rel < 1e-3 - assert final_rel < 5e-3 + assert final_rel < 2e-4 diff --git a/tests/test_sharded_segmentation.py b/tests/test_sharded_segmentation.py index a38ba0e..70a8aff 100644 --- a/tests/test_sharded_segmentation.py +++ b/tests/test_sharded_segmentation.py @@ -15,6 +15,7 @@ import mbirtorch import mbirtorch.preprocess as mtp +import mbirtorch.preprocess.mar as mtmar from mbirtorch import _sharding @@ -211,6 +212,115 @@ def test_sharded_bh_correction_matches_single_device(): assert rel < 1e-3 +# ── the plastic-coefficient floor, on one device with a view mask ──────────── +# This case is unsharded, but it exists only because of sharding: the view +# mask that reaches it is the padding indicator a multi-device placement +# builds. So it is gated here, beside the sharded MAR path it belongs to. +# The exponent tuples below are the ones bh_correction builds for one metal +# term with one cross term, which makes the two quantities inside +# _correct_plastic_sinogram simple enough to write out independently: +# Sp = theta[0] + theta[1] * m +# y_minus_Sm = clamp(y - theta[2] * m, min=0) +MASKED_FLOOR_EXPONENTS = [(1, 0), (1, 1), (0, 1)] +# Exact binary fractions, so Sp = theta[0] + theta[1] * m is an exact multiple +# of 1/512 (see the draw below) and its masked sum is exact in float32. +MASKED_FLOOR_THETA = np.array([0.25, 0.5, 0.25]) +# Sp runs over [0.25, 0.75) here, so this floor lands inside that range: it +# binds for most elements and leaves the rest to torch.maximum's other branch. +MASKED_FLOOR_GAMMA = 1.25 + + +def _masked_floor_case(num_views=6, real_views=4, det_shape=(5, 7)): + """A padded single-device sinogram triple plus its real-view mask. + + The draws are quantized to multiples of 1/256 so that the masked sum the + test turns on is EXACT in float32: every partial sum of Sp over the 140 + real pixels is a multiple of 1/512 below 128, which float32 holds exactly, + so no summation order and no device can move it. That is what makes the + two floors compared below -- and whether they differ -- a property of this + input rather than of the host that ran it. Drawing plain float32 leaves + the sum host-dependent, and the CUDA nightly read a sum whose two floors + rounded to the same float32, which left the last check with nothing to see. + """ + rng = np.random.default_rng(0) + full = (num_views,) + det_shape + + def draw(): + quantized = np.floor(rng.random(full) * 256.0) / 256.0 + array = torch.as_tensor(quantized.astype(np.float32)) + array[real_views:] = 0 # the padded views the engine zero-fills + return array + + plastic, metal, measured = draw(), draw(), draw() + view_mask = torch.as_tensor( + (np.arange(num_views) < real_views).reshape(num_views, 1, 1)) + num_real_pixels = real_views * det_shape[0] * det_shape[1] + return plastic, metal, measured, view_mask, num_real_pixels + + +def test_masked_single_device_plastic_floor_keeps_the_unsharded_arithmetic(): + """One plain tensor plus a view mask must take the float32 reduction. + + This is the branch a padded sinogram would reach if it were handed to + the correction as a single tensor, and it is the one combination the MAR + tests never drove. The sharded form sums each piece to a Python float + and divides in float64, which is a different rounding: the check below + pins the result to the float32 expression and, on the constructed input + above, shows the float64 form landing somewhere else -- so a change back + to it fails here instead of silently moving a single-device answer. + """ + plastic, metal, measured, view_mask, num_real_pixels = _masked_floor_case() + theta, gamma = MASKED_FLOOR_THETA, MASKED_FLOOR_GAMMA + + out = mtmar._correct_plastic_sinogram( + measured, plastic, [metal], theta, MASKED_FLOOR_EXPONENTS, + num_cross_terms=1, num_metal_terms=1, p_normalization=1.0, gamma=gamma, + view_mask=view_mask, num_real_pixels=num_real_pixels) + + plastic_coef = (torch.zeros_like(plastic) + + float(theta[0]) * torch.ones_like(plastic) + + float(theta[1]) * metal) + residual = torch.clamp(measured - float(theta[2]) * metal, min=0) + + # The expression this branch is required to keep: a float32 masked sum, + # divided by the real pixel count, held against Sp with torch.maximum. + float32_floor = gamma * (torch.sum(plastic_coef * view_mask) / float(num_real_pixels)) + expected = 1.0 * residual / torch.maximum(plastic_coef, float32_floor) + assert torch.equal(out, expected) + + # The float64 host divide, which is what the sharded branch must use and + # what this branch must not: on this input it moves the floor by one unit + # in the last place and changes most of the clamped elements. + float64_floor = gamma * (float(torch.sum(plastic_coef * view_mask)) / float(num_real_pixels)) + combined = 1.0 * residual / torch.clamp(plastic_coef, min=float64_floor) + + # Whether the two forms CAN differ at all, checked rather than assumed. + # Two things have to hold. The clamp receives the float64 floor as a + # float32 scalar, so the comparison that matters is between the two floors + # AS FLOAT32 -- a gap narrower than half a float32 step disappears in that + # cast and leaves nothing downstream to see (the CUDA nightly read a gap of + # 1.4e-08 against a step of 1.2e-07 and no element differed). And the + # floor has to bind somewhere with something to divide, since an element + # above both floors is divided by itself either way. The exact masked sum + # makes both facts host-independent, so the assert should hold wherever + # this runs; it is skipped with a message rather than failed if some + # machine still lands the two floors on one float32, because then the test + # has no discrimination to offer and would pass whichever form the library + # used. + floors_differ = float32_floor.item() != float(np.float32(float64_floor)) + binds = bool(((plastic_coef < float32_floor) & (residual > 0)).any()) + print("masked single-device Sp floor: " + f"float32 {float32_floor.item()!r} vs float64 {float64_floor!r}, " + f"{int((expected != combined).sum())} elements differ; " + f"floors differ as float32: {floors_differ}, floor binds: {binds}") + if floors_differ and binds: + assert not torch.equal(expected, combined) + else: + print('the two forms cannot be told apart on this input, so the check ' + 'above would pass either way; the inequality is skipped rather ' + 'than asserted') + + def test_sharded_save_and_export_stream_by_slab(tmp_path, monkeypatch): """Sharded saves gather one slab at a time (never the whole volume) and still write byte-identical files. The slab size is shrunk so several diff --git a/tests/test_sharding.py b/tests/test_sharding.py index ee5ae0f..a47bd67 100644 --- a/tests/test_sharding.py +++ b/tests/test_sharding.py @@ -53,6 +53,64 @@ def test_banded_adjoint_pair_values(): assert torch.allclose(total, partials[0] + partials[1]) +def test_the_streamed_reduce_matches_the_one_shot_sum_exactly(monkeypatch): + """The reduce moves each arriving partial in bounded row slabs, so the + owner never holds more than one slab per source above its running total. + + Streaming partitions the ELEMENTS: each element is still accumulated in + partial order, so the streamed result is bit for bit the one-shot sum and + not merely close to it. And the partials are read, never written, so a + caller may still use them afterwards. + """ + from mbirtorch import _sharding + owner = torch.device("cpu") + partials = [torch.rand(37, 5) for _ in range(4)] + untouched = [p.clone() for p in partials] + one_shot = ((partials[0] + partials[1]) + partials[2]) + partials[3] + # 40 bytes is two rows of a 5-column float32 band, so this runs 19 slabs + # rather than the single slab the default budget would give at this size. + monkeypatch.setattr(_sharding, "REDUCE_SLAB_BYTES", 40) + assert _sharding.reduce_slab_rows(37, 5 * 4) == 2 + total = sum_band_to_owner(partials, owner) + assert torch.equal(total, one_shot) + assert all(torch.equal(p, u) for p, u in zip(partials, untouched)) + # A single partial is still handed straight back, with no copy made. + assert sum_band_to_owner(partials[:1], owner) is partials[0] + + +def test_the_streamed_reduce_leaves_the_sharded_back_projection_unchanged( + monkeypatch): + """The test above pins the reduce; this one pins the driver that calls it. + + At test sizes a whole band fits inside one slab and moves in a single + piece, so the streaming path runs end to end only when the budget is + forced down. Without this the suite would never execute a multi-slab + reduce through the real back projection. + """ + import mbirtorch + from mbirtorch import _sharding + sino_shape = (9, 7, 8) # padded slices, 2 devices + angles = np.linspace(0, np.pi, sino_shape[0], endpoint=False) + + def build(devices): + m = mbirtorch.ParallelBeamModel(sino_shape, angles) + m.configure_devices(devices=["cpu"]) + m.set_params(no_warning=True, verbose=0) + if devices != ["cpu"]: + m.configure_devices(devices=devices) + return m + + rng = np.random.default_rng(11) + sino = rng.standard_normal(sino_shape).astype(np.float32) + reference = build(["cpu"]).back_project(sino) + # Two rows of a four-slice band at a time: many slabs, not one. + monkeypatch.setattr(_sharding, "REDUCE_SLAB_BYTES", 2 * 4 * 4) + m2 = build(["cpu", "cpu"]) + streamed = m2._gather_recon(m2.back_project(sino, output_sharded=True)) + rel = np.max(np.abs(streamed - reference)) / np.max(np.abs(reference)) + assert rel < 1e-5, rel + + def test_run_per_device_order_and_pool(): devs = ["cpu", "cpu", "cpu"] out = run_per_device(devs, lambda i, d: (i, str(d))) @@ -751,3 +809,851 @@ def cbuild(): crel = np.max(np.abs(cout - cref)) / max(np.max(np.abs(cref)), 1e-30) print(f"sparse-view cone n4 vs n1: rel {crel:.2e}") assert crel < 5e-3, crel # same calibration as the parallel case above + + +# ── the forward's column gather (default off) ──────────────────────────────── +# What may FAIL here, and what may only be recorded. The value bar these +# tests hold is the one the library already ships: the kernel-parity floor the +# suites above enforce, at the 1e-5 relative these cone cases use on CPU. The +# multi-GPU measurement also registered an EXPECTATION beside that floor -- the +# column gather sat about 1.5e-06 relative from the one-device anchor at the +# 1024-class cell (measured 2026-08-10 on four H100s, job mg10), against a +# banded walk that sat at its own repeat floor. That expectation is recorded +# so a later reading well outside it is visible to a human weighing the +# tradeoff; it is deliberately NOT a threshold, and nothing here asserts it. +# The distances below are printed for that comparison. On CPU the runs are +# deterministic and the gather's calls are the single-device call shape, so +# what these tests do assert is exact-path mechanics rather than that bar. +def _cone_column_case(devices, cell=(8, 8, 8), pixel_batch=None): + """A cone model on virtual CPU devices with the column gather switched + on, plus its single-device reference.""" + m, idx, vals, sino, ref_fwd, ref_back = _cone_banded_case(devices, cell) + m.forward_column_gather = True + if pixel_batch is not None: + m.forward_project_pixel_batch = pixel_batch + assert m._column_gather_forward() + return m, idx, vals, sino, ref_fwd, ref_back + + +def test_gather_column_band_assembles_the_full_height_cylinder(): + # The primitive: every slice-owner's rows [p0:p1] moved to one target and + # concatenated along the SLICE axis, in shard (global slice) order. A + # single shard short-circuits the concatenation. + from mbirtorch._sharding import gather_column_band + rng = np.random.default_rng(11) + full = torch.as_tensor(rng.standard_normal((9, 6)).astype(np.float32)) + shards = [full[:, 0:2].contiguous(), full[:, 2:4].contiguous(), + full[:, 4:6].contiguous()] + cyl = gather_column_band(shards, 3, 7, torch.device("cpu")) + assert cyl.shape == (4, 6) + assert torch.equal(cyl, full[3:7]) + # A degenerate range is legal and empty; one shard is returned as itself. + empty = gather_column_band(shards, 5, 5, torch.device("cpu")) + assert empty.shape == (0, 6) + one = gather_column_band(shards[:1], 0, 9, torch.device("cpu")) + assert torch.equal(one, shards[0]) + # The host-bounce path is value-correct too (dev2dev_safe False). + bounced = gather_column_band(shards, 0, 9, torch.device("cpu"), + dev2dev_safe=False) + assert torch.equal(bounced, full) + + +@pytest.mark.skipif(not torch.backends.mps.is_available(), + reason="needs a second local device (mps)") +def test_gather_column_band_moves_across_real_devices(): + from mbirtorch._sharding import gather_column_band + full = torch.rand(32, 8) + shards = [full[:, :4].contiguous().to("cpu"), + full[:, 4:].contiguous().to("mps")] + cyl = gather_column_band(shards, 8, 16, torch.device("mps")) + assert cyl.device.type == "mps" + assert torch.allclose(cyl.cpu(), full[8:16], atol=1e-6) + + +def test_column_gather_matches_single_device_at_every_batch(monkeypatch): + # The values gate on virtual CPU devices: a full-height call at + # slice_start=0 is the single-device call shape, so the gathered forward + # must reproduce the single-device values -- at one batch covering the + # pass, and at batches that force several. The environment is cleared + # first so a suite run forcing the banded walk cannot unseat the gather + # this test is about -- the same pinning the banded tests do in reverse. + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + for n in (2, 3): + for batch in (None, 1, 5, 10 ** 6): + m, idx, vals, _sino, ref_fwd, _ref_back = _cone_column_case( + ["cpu"] * n, pixel_batch=batch) + fwd = m._gather_sinogram(m.sparse_forward_project(vals, idx)) + rel = np.max(np.abs(fwd - ref_fwd)) / np.max(np.abs(ref_fwd)) + print(f"cone column gather n={n} batch={batch}: rel {rel:.2e}") + assert rel < 1e-5, (n, batch, rel) + + +def test_column_gather_holds_the_adjoint_and_the_padded_forms(monkeypatch): + # The back driver is untouched, so the pair must stay adjoint with the + # gather on -- on a padded cell (9 views, 7 rows over 2 devices pads both + # axes), where the gathered cylinder carries the inert padded slice tail. + # The environment is cleared first, for the reason above. + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + m, idx, vals, sino, ref_fwd, ref_back = _cone_column_case( + ["cpu", "cpu"], cell=(9, 7, 8), pixel_batch=4) + assert m.recon_placement.is_padded and m.sino_placement.is_padded + fwd = m.sparse_forward_project(vals, idx) + back = m.sparse_back_project(sino, idx) + real = vals.shape[1] + assert np.allclose(m._gather_sinogram(fwd), ref_fwd, atol=1e-5) + assert np.allclose(back.gather()[:, :real], ref_back, atol=1e-5) + lhs = float(np.sum(m._gather_sinogram(fwd) * sino)) + rhs = float(np.sum(vals * back.gather()[:, :real])) + assert abs(lhs - rhs) / max(abs(rhs), 1e-30) < 1e-4, (lhs, rhs) + + +def test_column_gather_replaces_the_band_broadcast(monkeypatch): + # The mechanics witness. With the gather on, the cone forward must call + # gather_column_band and must NOT broadcast a band; each gather takes one + # piece per slice-owner and yields a cylinder that is the batch wide and + # the WHOLE device-form slice axis tall; and each projector call runs at + # slice_start=0 over that whole axis for the owner's own views. The + # environment is cleared first, for the reason above. + from mbirtorch import _sharding as sharding + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + batch, n = 4, 2 + m, idx, vals, _sino, _ref_fwd, _ref_back = _cone_column_case( + ["cpu"] * n, pixel_batch=batch) + slices = m.recon_placement.padded_size + gathers, broadcasts, calls = [], [], [] + real_gather = sharding.gather_column_band + + def spy_gather(shard_tensors, p0, p1, target, dev2dev_safe=True): + out = real_gather(shard_tensors, p0, p1, target, dev2dev_safe) + gathers.append((len(shard_tensors), p0, p1, tuple(out.shape))) + return out + + def spy_broadcast(*args, **kwargs): + broadcasts.append(args) + raise AssertionError("the column gather must not broadcast a band") + + real_call = m.projector_functions.sparse_forward_project_view_range + + def spy_call(band_values, pixel_indices, view_range, slice_start=0, + dev_index=0, plan=None, accumulate_into=None): + calls.append((tuple(band_values.shape), int(pixel_indices.shape[0]), + tuple(view_range), slice_start)) + return real_call(band_values, pixel_indices, view_range, + slice_start=slice_start, dev_index=dev_index, + plan=plan, accumulate_into=accumulate_into) + + monkeypatch.setattr(sharding, "gather_column_band", spy_gather) + monkeypatch.setattr(sharding, "broadcast_band_to_views", spy_broadcast) + monkeypatch.setattr(m.projector_functions, + "sparse_forward_project_view_range", spy_call) + m.sparse_forward_project(vals, idx) + + expected_batches = -(-len(idx) // batch) + assert not broadcasts + assert len(gathers) == n * expected_batches + for pieces, p0, p1, shape in gathers: + assert pieces == n # one piece per slice-owner + assert shape == (p1 - p0, slices) # the batch, at every slice + assert p1 - p0 <= batch + # One projector call per (pixel batch, view-owner), each over the whole + # slice axis anchored at 0 and over that owner's own real views. + assert len(calls) == n * expected_batches + spans = [(v0, v1) for _, _, (v0, v1), _ in calls] + for cyl_shape, n_pixels, (v0, v1), slice_start in calls: + assert slice_start == 0 and cyl_shape[1] == slices + assert n_pixels == cyl_shape[0] and v1 > v0 + assert set(spans) == { + (v0, v0 + valid) for _d, (v0, _v1), valid + in m.sino_placement.padded_shard_ranges() if valid > 0} + + +def test_the_column_gather_runs_one_batch_ahead_of_the_projection(monkeypatch): + # The prefetch witness. Each view-owner issues the NEXT pixel batch's + # gather before it projects the current batch, so that on real devices the + # copies feeding one projection can be moving while another projection + # runs. On virtual CPU devices nothing moves and nothing can be timed, so + # what is asserted here is the ORDER the driver issues its work in, which + # is the part of the change that has to hold on every device. + # + # The order one worker records is g0, g1, p0, g2, p1, ... , g(K-1), p(K-2), + # p(K-1): batch k+1's gather is issued before batch k is projected, the + # first gather runs before the loop, and the last batch has nothing to + # gather ahead of it. The entry and exit of each projection are both + # recorded, so the witness is not merely that the gather precedes the + # accumulation -- it precedes the projector call entirely. + # + # Each worker runs in its own thread, so events are kept per thread. A + # pool thread is allowed to run more than one worker when one finishes + # before the next is submitted, and it would then record two workers' + # sequences end to end; the check reads blocks rather than the whole list + # so that it witnesses the order either way. + # + # The environment knob is cleared first, the way the tests around this one + # already do: it forces the gather off whatever the model says, so a suite + # run that sets it would otherwise decide what this test measures. + import threading + from mbirtorch import _sharding as sharding + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + batch, n = 4, 2 + m, idx, vals, _sino, ref_fwd, _rb = _cone_column_case( + ["cpu"] * n, pixel_batch=batch) + n_batches = -(-len(idx) // batch) + assert n_batches > 1 # or there is no prefetch to see + events = {} + real_gather = sharding.gather_column_band + real_call = m.projector_functions.sparse_forward_project_view_range + + def spy_gather(shard_tensors, p0, p1, target, dev2dev_safe=True): + # Recorded at ENTRY: what is being witnessed is when the gather is + # issued, not when it returns. + events.setdefault(threading.get_ident(), []).append(f'g{p0 // batch}') + return real_gather(shard_tensors, p0, p1, target, dev2dev_safe) + + def spy_call(band_values, pixel_indices, view_range, slice_start=0, + dev_index=0, plan=None, accumulate_into=None): + seq = events.setdefault(threading.get_ident(), []) + # Number the projections within THIS worker, which begins at its own + # first gather, so that a thread running a second worker starts over + # at zero rather than counting on from the first. + first = len(seq) - 1 - seq[::-1].index('g0') + k = sum(1 for e in seq[first:] if e.endswith('-in')) + seq.append(f'p{k}-in') + block = real_call(band_values, pixel_indices, view_range, + slice_start=slice_start, dev_index=dev_index, + plan=plan, accumulate_into=accumulate_into) + seq.append(f'p{k}-out') + return block + + monkeypatch.setattr(sharding, "gather_column_band", spy_gather) + monkeypatch.setattr(m.projector_functions, + "sparse_forward_project_view_range", spy_call) + fwd = m._gather_sinogram(m.sparse_forward_project(vals, idx)) + + expected = ['g0'] + for k in range(n_batches): + if k + 1 < n_batches: + expected.append(f'g{k + 1}') + expected += [f'p{k}-in', f'p{k}-out'] + assert events, "the column gather did not run" + recorded = 0 + for seq in events.values(): + # Every worker of this cell owns real views, so each ran the whole + # sequence; a thread holds a whole number of them. + assert len(seq) % len(expected) == 0, seq + for start in range(0, len(seq), len(expected)): + assert seq[start:start + len(expected)] == expected, seq + recorded += 1 + assert recorded == n # one sequence per view-owner + # The prefetch moves WHEN a gather is issued and nothing else, so the + # values are the ones the path already produced. + assert np.allclose(fwd, ref_fwd, atol=1e-5) + + # And the values hold across batch widths that force several batches, + # including one that leaves a short final batch (30 pixels over 7). + for width in (1, 3, 7): + mb, idxb, valsb, _s, ref_b, _rb2 = _cone_column_case( + ["cpu"] * n, pixel_batch=width) + out = mb._gather_sinogram(mb.sparse_forward_project(valsb, idxb)) + rel = np.max(np.abs(out - ref_b)) / np.max(np.abs(ref_b)) + print(f"cone gather one batch ahead, {width}-pixel batches: " + f"rel {rel:.2e}") + assert rel < 1e-5, (width, rel) + + +def test_the_column_gather_is_on_by_default_and_scoped_to_its_geometry( + monkeypatch): + # The switch: on unless refused, refused on a geometry the shape has never + # been measured on however it is asked, and overridable from the + # environment either way so one session can run both shapes over the same + # inputs. The environment is cleared first, because this test reads the + # DEFAULT and a suite run may be forcing the path around it. + import mbirtorch + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + cone, _idx, _vals, _sino, _f, _b = _cone_banded_case(["cpu", "cpu"]) + assert cone.column_gather_geometry + assert cone._column_gather_forward() # default on + cone.forward_column_gather = False # the rollback + assert not cone._column_gather_forward() + cone.forward_column_gather = True + assert cone._column_gather_forward() + + # The row-aligned geometry declares the same capability, on its own + # measurement, and is on by default in the same way. + angles = np.linspace(0, np.pi, 8, endpoint=False) + par = mbirtorch.ParallelBeamModel((8, 6, 8), angles) + assert par.column_gather_geometry + assert par._column_gather_forward() + par.forward_column_gather = False + assert not par._column_gather_forward() + par.forward_column_gather = True + assert par._column_gather_forward() + + # A geometry that has never been timed on the shape refuses it however it + # is asked: translation shares cone's banded branch, and an argument that + # it should gain too is not a measurement. + trans = mbirtorch.TranslationModel( + (4, 6, 8), np.zeros((4, 3), dtype=np.float32), + source_detector_dist=32.0, source_iso_dist=16.0) + trans.forward_column_gather = True + assert not trans.column_gather_geometry + assert not trans._column_gather_forward() + + # The environment wins over the attribute in BOTH directions: `cone` holds + # an explicit True and `refused` an explicit False, and each env value + # drives the two models to the same answer. + import os + refused = _cone_banded_case(["cpu", "cpu"])[0] + refused.forward_column_gather = False + for value, expected in (("1", True), ("on", True), + ("0", False), ("off", False)): + os.environ[COLUMN_GATHER_ENV_VAR] = value + try: + assert refused._column_gather_forward() is expected + assert cone._column_gather_forward() is expected + finally: + del os.environ[COLUMN_GATHER_ENV_VAR] + assert not refused._column_gather_forward() + unset = _cone_banded_case(["cpu", "cpu"])[0] + assert unset._column_gather_forward() # the shipped default + + +@pytest.mark.parametrize('geometry', ('cone', 'parallel')) +def test_the_banded_walk_is_what_runs_with_the_switch_off(geometry, + monkeypatch): + # The rollback, exercised on both geometries that can take the gather: + # switching the gather off selects the banded walk, which broadcasts + # bands and gathers no columns. The switch has to be refused explicitly + # now that unset means on, and the environment knob is cleared first for + # the reason above. + from mbirtorch import _sharding as sharding + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + if geometry == 'cone': + m, idx, vals, _sino, ref_fwd, _rb = _cone_banded_case(["cpu", "cpu"]) + else: + m, idx, vals, _sino, ref_fwd, _rb, _b2 = _banded_case(["cpu", "cpu"]) + m.forward_column_gather = False + assert m.column_gather_geometry and not m._column_gather_forward() + broadcasts = [] + real_broadcast = sharding.broadcast_band_to_views + + def spy_broadcast(band, view_owners, dev2dev_safe=True): + broadcasts.append(tuple(band.shape)) + return real_broadcast(band, view_owners, dev2dev_safe) + + def refuse(*args, **kwargs): + raise AssertionError("the banded walk must not gather columns") + + monkeypatch.setattr(sharding, "broadcast_band_to_views", spy_broadcast) + monkeypatch.setattr(sharding, "gather_column_band", refuse) + fwd = m._gather_sinogram(m.sparse_forward_project(vals, idx)) + assert broadcasts + assert np.allclose(fwd, ref_fwd, atol=1e-5) + + +def test_column_gather_recon_matches_single_device(monkeypatch): + # The end-to-end gate: a seeded cone reconstruction on two virtual CPU + # devices with the gather on must reproduce the single-device run, which + # is where the two changed summation orders (the vertical sum into the + # body, the pixel sum out of it) would show up if they were not inside + # the value class the forward already has. The environment is cleared + # first so each of the three runs below is the shape it names. + import mbirtorch + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + cell = (8, 8, 8) + angles = np.linspace(0, 2 * np.pi, cell[0], endpoint=False) + + def build(devices): + m = mbirtorch.ConeBeamModel(cell, angles, source_detector_dist=32, + source_iso_dist=16) + m.configure_devices(devices=["cpu"]) + m.set_params(no_warning=True, verbose=0) + if len(devices) > 1: + m.configure_devices(devices=devices) + return m + + m1 = build(["cpu"]) + rs = tuple(m1.get_params('recon_shape')) + phantom = np.zeros(rs, dtype=np.float32) + phantom[1:-1, 1:-1, 1:-1] = 1.0 + sino = m1.forward_project(phantom) + np.random.seed(31) + ref, _ = m1.recon(sino, max_iterations=2, stop_threshold_change_pct=0.0) + + banded = build(["cpu", "cpu"]) + banded.forward_column_gather = False # unset now means the gather + np.random.seed(31) + banded_out, _ = banded.recon(sino, max_iterations=2, + stop_threshold_change_pct=0.0) + gathered = build(["cpu", "cpu"]) + gathered.forward_column_gather = True + gathered.forward_project_pixel_batch = 8 + np.random.seed(31) + out, _ = gathered.recon(sino, max_iterations=2, + stop_threshold_change_pct=0.0) + scale = max(np.max(np.abs(ref)), 1e-30) + rel = np.max(np.abs(out - ref)) / scale + rel_banded = np.max(np.abs(banded_out - ref)) / scale + # Printed rather than asserted against each other: which of the two sits + # closer to the anchor is the reading the registered expectation is for. + print(f"cone recon vs n1: column gather {rel:.2e}, " + f"banded {rel_banded:.2e}") + assert rel < 5e-3, rel # the shipped parity floor, as above + + +# ── the same gather, on the row-aligned geometry ───────────────────────────── +# Parallel takes the column gather for a different measured reason than cone. +# It CAN produce its detector rows from a slice band -- the banded walk does +# exactly that -- but its forward kernel runs about twice as efficiently per +# slice on the full-width block of values the gather hands it as on the +# shard-width blocks the band hands it (measured 2026-08-10 on one H100, at +# 0.0411 ms per slice on a 1008-wide block against 0.0823 on a 504-wide one, +# with the device count held at one). +# +# The value bar was expected to be EQUALITY here, on the argument that each +# detector row keeps a single producing call and CPU sums are deterministic. +# The row half of that is true, and the mechanics test below asserts it +# directly. Equality is not, and the measurement that settled it is recorded +# because it is worth knowing before anyone tries again (2026-08-10, this +# suite, virtual CPU devices). Run first in a fresh interpreter, BOTH the +# banded walk and the column gather reproduce the single-device sinogram bit +# for bit. Run once other shapes have gone through the same per-device +# bodies -- which is what a full suite run does -- both land in the float32 +# epsilon class instead, the banded walk at 1.6e-07 to 4.0e-07 and the gather +# at 1.1e-07 to 4.0e-07 over the same cells. The cause is the same for both: +# the per-device bodies are separately torch.compiled, and what a compiled +# body emits depends on the shapes its instance has already seen, so two +# devices can differ in the last bit on identical inputs. Bit-equality is +# therefore a property of the process, not of the driver shape, and these +# tests hold the gather against the shape it replaces instead. +def _parallel_column_case(devices, sino_shape=(8, 6, 8), pixel_batch=None): + """A parallel model on virtual CPU devices with the column gather + switched on, plus its single-device reference.""" + m, idx, vals, sino, ref_fwd, ref_back, _b2 = _banded_case(devices, + sino_shape) + m.forward_column_gather = True + if pixel_batch is not None: + m.forward_project_pixel_batch = pixel_batch + assert m._column_gather_forward() + return m, idx, vals, sino, ref_fwd, ref_back + + +def test_parallel_column_gather_matches_the_shape_it_replaces(monkeypatch): + # The environment is cleared first so each leg below runs the shape it + # names, whatever a suite run is forcing around this test. + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + # The values gate. Both shapes are run over the same inputs in the same + # process and both are held to the same bar, which is the reading that + # does not move with the compile state above; the two distances are + # printed side by side for the same reason. One call per view-owner over + # every pixel is the single-device call in every respect that sets a + # value: the same voxel columns in one array, the whole slice range + # anchored at 0, and each detector row produced by that one call and no + # other. A row taking contributions from more than one call would show up + # here as an order-one error, not as a last bit. + for n in (2, 3): + # None takes the shipped batch, which covers a pass this size in one + # call; the large value asks for that explicitly. + for batch in (None, 10 ** 6): + m, idx, vals, _sino, ref_fwd, _rb = _banded_case(["cpu"] * n)[:6] + m.forward_column_gather = False # unset now means the gather + banded = m._gather_sinogram(m.sparse_forward_project(vals, idx)) + m.forward_column_gather = True + if batch is not None: + m.forward_project_pixel_batch = batch + fwd = m._gather_sinogram(m.sparse_forward_project(vals, idx)) + scale = np.max(np.abs(ref_fwd)) + rel = np.max(np.abs(fwd - ref_fwd)) / scale + rel_banded = np.max(np.abs(banded - ref_fwd)) / scale + print(f"parallel column gather n={n} batch={batch}: rel {rel:.2e}," + f" banded {rel_banded:.2e}") + assert rel < 1e-5 and rel_banded < 1e-5, (n, batch, rel, + rel_banded) + + # The one summation order this shape does change for a row-aligned + # geometry: several pixel batches turn a single accumulation over every + # pixel into a host-side sum of per-batch partials. Nothing about the + # rows moves, so what is left is float noise in the same class as above + # (measured 1.0e-07 to 1.6e-07 here). This is the case that runs at + # production sizes, where the pass is far wider than one batch. + for batch in (1, 5, 7): + m, idx, vals, _sino, ref_fwd, _rb = _parallel_column_case( + ["cpu", "cpu"], pixel_batch=batch) + fwd = m._gather_sinogram(m.sparse_forward_project(vals, idx)) + rel = np.max(np.abs(fwd - ref_fwd)) / np.max(np.abs(ref_fwd)) + print(f"parallel column gather, {batch}-pixel batches: rel {rel:.2e}") + assert rel < 1e-5, (batch, rel) + + +def test_parallel_column_gather_gathers_columns_and_sizes_its_rows_by_them( + monkeypatch): + # The environment is cleared first, for the reason above. + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + # The mechanics witness, plus the row-aligned fact the banded walk used to + # supply by construction. With the gather on, the parallel forward calls + # gather_column_band and broadcasts no band; each cylinder is one pixel + # batch by the WHOLE device-form slice axis; each projector call runs at + # slice_start=0 over that whole axis for the owner's own views; and the + # block that comes back is as TALL as the cylinder, because a row-aligned + # body sizes its output by the values it was handed. That last one is why + # the assembled shard carries the device form's padded row count and not + # the real detector rows. + from mbirtorch import _sharding as sharding + batch, n = 4, 2 + m, idx, vals, _sino, _ref_fwd, _rb = _parallel_column_case( + ["cpu"] * n, pixel_batch=batch) + slices = m.recon_placement.padded_size + channels = int(m.get_params('sinogram_shape')[2]) + gathers, calls = [], [] + real_gather = sharding.gather_column_band + + def spy_gather(shard_tensors, p0, p1, target, dev2dev_safe=True): + out = real_gather(shard_tensors, p0, p1, target, dev2dev_safe) + gathers.append((len(shard_tensors), p0, p1, tuple(out.shape))) + return out + + def spy_broadcast(*args, **kwargs): + raise AssertionError("the column gather must not broadcast a band") + + real_call = m.projector_functions.sparse_forward_project_view_range + + def spy_call(band_values, pixel_indices, view_range, slice_start=0, + dev_index=0, plan=None, accumulate_into=None): + block = real_call(band_values, pixel_indices, view_range, + slice_start=slice_start, dev_index=dev_index, + plan=plan, accumulate_into=accumulate_into) + calls.append((tuple(band_values.shape), tuple(view_range), slice_start, + tuple(block.shape))) + return block + + monkeypatch.setattr(sharding, "gather_column_band", spy_gather) + monkeypatch.setattr(sharding, "broadcast_band_to_views", spy_broadcast) + monkeypatch.setattr(m.projector_functions, + "sparse_forward_project_view_range", spy_call) + fwd = m.sparse_forward_project(vals, idx) + + expected_batches = -(-len(idx) // batch) + assert len(gathers) == n * expected_batches + for pieces, p0, p1, shape in gathers: + assert pieces == n # one piece per slice-owner + assert shape == (p1 - p0, slices) # the batch, at every slice + assert p1 - p0 <= batch + assert len(calls) == n * expected_batches + for cyl_shape, (v0, v1), slice_start, block_shape in calls: + assert slice_start == 0 and cyl_shape[1] == slices + assert block_shape == (v1 - v0, slices, channels) + assert set((v0, v1) for _c, (v0, v1), _s, _b in calls) == { + (v0, v0 + valid) for _d, (v0, _v1), valid + in m.sino_placement.padded_shard_ranges() if valid > 0} + # And the shard the driver assembles carries those same rows. + assert all(tuple(t.shape[1:]) == (slices, channels) for t in fwd.tensors) + + +def test_parallel_column_gather_holds_the_padded_and_sparse_view_forms( + monkeypatch): + # The environment is cleared first, for the reason above. + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + # The two forms where a row-aligned geometry's DEVICE shape differs from + # its problem shape. A padded slice axis pads the sinogram's detector + # rows with it, so every block this driver assembles -- including the + # empty one it builds for a view-owner with no real views -- has to be the + # padded row count. Sized at the real detector rows instead, which is the + # count a row-RANGE geometry's blocks carry, the shards do not concatenate + # at all. + for shape, devs in (((9, 7, 8), 2), # both axes padded + ((3, 7, 8), 4)): # padded rows, an empty owner + m, idx, vals, sino, ref_fwd, ref_back = _parallel_column_case( + ["cpu"] * devs, sino_shape=shape, pixel_batch=10 ** 6) + assert m.recon_placement.is_padded and m.sino_placement.is_padded + real_rows = shape[1] + fwd = m.sparse_forward_project(vals, idx) + assert all(t.shape[1] == m.recon_placement.padded_size + for t in fwd.tensors), shape + # The padded row tail stays identically zero, as the entry fill left + # it: the gathered cylinder's padded slice tail is zero, and a + # row-aligned body maps those columns straight to those rows. + assert max(float(t[:, real_rows:].abs().max()) for t in fwd.tensors) \ + == 0.0, shape + assert np.allclose(m._gather_sinogram(fwd), ref_fwd, atol=1e-5), shape + # The back driver is untouched, so the pair stays adjoint. + back = m.sparse_back_project(sino, idx) + assert np.allclose(back.gather()[:, :vals.shape[1]], ref_back, + atol=1e-5) + lhs = float(np.sum(m._gather_sinogram(fwd) * sino)) + rhs = float(np.sum(vals * back.gather()[:, :vals.shape[1]])) + assert abs(lhs - rhs) / max(abs(rhs), 1e-30) < 1e-4, (shape, lhs, rhs) + + +def test_column_gather_batch_accumulation_matches_the_shape_it_replaces( + monkeypatch): + # The environment is cleared first, for the reason above. + from mbirtorch.tomography_model import COLUMN_GATHER_ENV_VAR + monkeypatch.delenv(COLUMN_GATHER_ENV_VAR, raising=False) + # Each pixel batch after the first adds into the owner's block from inside + # the projector's view loop, rather than assembling its own block for the + # driver to add afterwards. Those are the same summands added in the same + # order, element for element, so the bar here is EQUALITY and not closeness. + # + # This one is safe to assert bit for bit whatever the compile state, unlike + # the cross-device comparisons above. Both legs drive the SAME per-device + # compiled bodies over the SAME shapes in the SAME process, so every block + # entering the accumulation is identical by construction and the legs differ + # only in the arithmetic that combines them. The caveat recorded above is + # about two DEVICES emitting different code for one shape, which cannot + # separate two legs that share their devices. + # + # There IS a second thing that separates two runs, and it has to be held + # still for the equality above to mean anything: torch's CPU scatter reduces + # in PARALLEL, so the body is not reproducible run to run once the problem + # is big enough to thread -- one shape run twice already differs from + # itself. Measured 2026-08-11 in a full suite run, this cell at two devices + # and 5-pixel batches: one shape against itself 5.2e-08, and the two shapes + # against each other 1.0e-07, which is that same noise drawn again and then + # carried through eight batches of accumulation rather than any reordering. + # On a 64x48x64 cell over 4096 pixels at 10 threads all three comparisons + # sat at 7.5e-08 together, the change adding nothing over the noise. + # + # So the threads are pinned to one below. That removes the only thing that + # separates two runs of the same arithmetic and lets this test assert what + # it is actually about -- that moving the addition does not move the + # values -- rather than measuring the scatter's thread scheduling. Pinned, + # every case here is bit-equal, including the ones that are not when the + # scatter is free to thread. + + def prior_shape(real_call, accumulating): + """The accumulation as it stood before it moved into the view loop: + every call assembles a block of its own, and the running block is added + to it afterwards. Counts the calls that were asked to accumulate, so + the comparison below cannot pass by never exercising the new arm.""" + def call(band_values, pixel_indices, view_range, slice_start=0, + dev_index=0, plan=None, accumulate_into=None): + block = real_call(band_values, pixel_indices, view_range, + slice_start=slice_start, dev_index=dev_index, + plan=plan) + if accumulate_into is None: + return block + accumulating.append(1) + accumulate_into.add_(block) + return accumulate_into + return call + + # Both geometries, two and three virtual CPU devices, and batches small + # enough that the pass runs many of them -- which is the case the fusion + # exists for and the only one where the two shapes can differ at all. + threads = torch.get_num_threads() + torch.set_num_threads(1) + try: + for name, case in (("parallel", _parallel_column_case), + ("cone", _cone_column_case)): + for n in (2, 3): + for batch in (1, 3, 5): + m, idx, vals = case(["cpu"] * n, pixel_batch=batch)[:3] + batches = -(-len(idx) // batch) + assert batches >= 2, (name, batch) + fused = np.asarray( + m._gather_sinogram(m.sparse_forward_project(vals, idx))) + # The same shape run twice, as the control: with the threads + # pinned this is exact, and a case where it were not would + # mean the noise above had another source and the comparison + # below could not be read as an ordering test. + control = np.asarray( + m._gather_sinogram(m.sparse_forward_project(vals, idx))) + assert np.array_equal(fused, control), (name, n, batch) + real_call = (m.projector_functions + .sparse_forward_project_view_range) + accumulating = [] + with monkeypatch.context() as mp: + mp.setattr(m.projector_functions, + "sparse_forward_project_view_range", + prior_shape(real_call, accumulating)) + prior = np.asarray(m._gather_sinogram( + m.sparse_forward_project(vals, idx))) + assert np.array_equal(fused, prior), (name, n, batch) + # Every view-owner with real views accumulates on all but + # its first batch, so the new arm ran once per (owner, + # batch) less one batch per owner. + owners = sum(1 for _d, (_v0, _v1), valid + in m.sino_placement.padded_shard_ranges() + if valid > 0) + assert len(accumulating) == owners * (batches - 1), ( + name, n, batch, len(accumulating), owners, batches) + + # The parameter itself, at the projector: handed a block it adds into + # that block and hands back the same object; handed None it allocates + # and writes. Accumulating one call's values onto another's therefore + # doubles them exactly. Inside the pinned region with the rest, because + # this compares two separate evaluations of the same body and the free + # scatter separates those on its own. + m, idx, vals = _banded_case(["cpu"])[:3] + pf = m.projector_functions + num_views = int(m.get_params('sinogram_shape')[0]) + t_vals = torch.as_tensor(vals) + t_idx = torch.as_tensor(idx, dtype=torch.int64) + once = pf.sparse_forward_project_view_range(t_vals, t_idx, + (0, num_views)) + running = pf.sparse_forward_project_view_range(t_vals, t_idx, + (0, num_views)) + assert torch.equal(once, running) # the control, as above + same = pf.sparse_forward_project_view_range(t_vals, t_idx, + (0, num_views), + accumulate_into=running) + assert same is running + assert torch.equal(running, once + once) + finally: + torch.set_num_threads(threads) + + +def test_parallel_column_gather_recon_matches_single_device(): + # The end-to-end gate, where the subset passes call the forward on small + # pixel sets and the pixel batch above therefore bites: a seeded parallel + # reconstruction on two virtual CPU devices with the gather on must + # reproduce the single-device run within the loop's own multi-device + # floor, which the banded walk beside it is read against. + import mbirtorch + sino_shape = (8, 6, 8) + angles = np.linspace(0, np.pi, sino_shape[0], endpoint=False) + + def build(devices): + m = mbirtorch.ParallelBeamModel(sino_shape, angles) + m.configure_devices(devices=["cpu"]) + m.set_params(no_warning=True, verbose=0) + if len(devices) > 1: + m.configure_devices(devices=devices) + return m + + m1 = build(["cpu"]) + rs = tuple(m1.get_params('recon_shape')) + phantom = np.zeros(rs, dtype=np.float32) + phantom[1:-1, 1:-1, 1:-1] = 1.0 + sino = m1.forward_project(phantom) + np.random.seed(31) + ref, _ = m1.recon(sino, max_iterations=3, stop_threshold_change_pct=0.0) + + banded = build(["cpu", "cpu"]) + np.random.seed(31) + banded_out, _ = banded.recon(sino, max_iterations=3, + stop_threshold_change_pct=0.0) + gathered = build(["cpu", "cpu"]) + gathered.forward_column_gather = True + gathered.forward_project_pixel_batch = 8 + np.random.seed(31) + out, _ = gathered.recon(sino, max_iterations=3, + stop_threshold_change_pct=0.0) + scale = max(np.max(np.abs(ref)), 1e-30) + rel = np.max(np.abs(out - ref)) / scale + rel_banded = np.max(np.abs(banded_out - ref)) / scale + print(f"parallel recon vs n1: column gather {rel:.2e}, " + f"banded {rel_banded:.2e}") + assert rel < 5e-4, rel # the sharded VCD loop's own floor at this cell + + +# ── one pixel at a time ────────────────────────────────────────────────────── +# The column gather's pixel batching hands the projectors a one-pixel call +# whenever a batch, or the remainder of a batch, is a single pixel, and a user +# can ask for one directly. On linux with torch 2.13.0, CPU inductor +# miscompiles exactly that case in both parallel bodies and lands the pixel's +# footprint one detector channel off (measured 2026-08-11: 6.56e-02 relative on +# the forward, 5.04e-02 on the back, on the 8x6x8 cell below; eager is right, +# and so is every width of two or more). The driver pads a one-pixel call to +# two and takes the padding back out. These two tests hold that: the first is +# the property a user cares about, the second is the padding itself. Both pass +# on any machine whose compiler is sound -- macOS is one -- so their value is +# the linux nightly. +def test_parallel_solo_pixel_projections_match_the_full_pass(): + # A pixel projects the same whether it is asked for alone or with the + # others. Forward: the projections of the single pixels sum to the whole + # pass, because the forward is linear in the voxels and each pixel writes + # its own footprint into the same sinogram. Back: one pixel's cylinder is + # that pixel's row of the whole pass, computed from the same sinogram. A + # body that reads a one-pixel call differently shows up here as an + # order-one error, not as a last bit. + m, idx, vals, sino, ref_fwd, ref_back, _b2 = _banded_case(["cpu"]) + solo_fwd = np.zeros_like(ref_fwd) + for i in range(len(idx)): + solo_fwd += m.sparse_forward_project(vals[i:i + 1], + idx[i:i + 1]).cpu().numpy() + rel = np.max(np.abs(solo_fwd - ref_fwd)) / np.max(np.abs(ref_fwd)) + print(f"parallel solo-pixel forward sum: rel {rel:.2e}") + assert rel < 1e-5, rel + + for i in (0, 1, len(idx) // 2, len(idx) - 1): + row = m.sparse_back_project(sino, idx[i:i + 1]).cpu().numpy() + assert row.shape == (1, ref_back.shape[1]) + rel_back = (np.max(np.abs(row[0] - ref_back[i])) + / np.max(np.abs(ref_back[i]))) + print(f"parallel solo-pixel back, pixel {i}: rel {rel_back:.2e}") + assert rel_back < 1e-5, (i, rel_back) + + +def test_the_minimum_pixel_width_padding_keeps_the_values(): + # The padding itself, held against the eager bodies it must agree with. + # The forward's padded column carries zero values at a repeated pixel + # index, and the forward output has no pixel axis, so the padded call is + # bit-identical and nothing is sliced off. The back's output does carry + # the pixel axis, so the padded call's extra row is sliced away and the row + # that stays must be exact, not close -- at both coefficient powers. A + # call that is already wide enough goes through untouched. + from mbirtorch import ConeBeamModel, projectors + from mbirtorch.parallel_beam import (ParallelBeamModel, + _parallel_back_view_batch, + _parallel_forward_view_batch) + # Declared by the geometry whose bodies need it, and by no other. + assert ParallelBeamModel.min_compiled_pixel_width == 2 + assert ConeBeamModel.min_compiled_pixel_width == 1 + + m, idx, vals, sino, _rf, _rb, _b2 = _banded_case(["cpu"]) + args = m._view_batch_args() + view_params = torch.as_tensor(np.asarray(m.get_params('angles')), + dtype=torch.float32) + one_idx = torch.as_tensor(idx[7:8], dtype=torch.int64) + one_vals = torch.as_tensor(vals[7:8]) + sino_t = torch.as_tensor(sino) + widths = [] + + def spy_forward(values, pixel_indices, *a, **kw): + widths.append(int(pixel_indices.shape[0])) + return _parallel_forward_view_batch(values, pixel_indices, *a, **kw) + + def spy_back(sino_batch, pixel_indices, *a, **kw): + widths.append(int(pixel_indices.shape[0])) + return _parallel_back_view_batch(sino_batch, pixel_indices, *a, **kw) + + padded_fwd = projectors.forward_at_min_pixel_width(spy_forward, 2) + padded_back = projectors.back_at_min_pixel_width(spy_back, 2) + + assert torch.equal( + padded_fwd(one_vals, one_idx, view_params, **args), + _parallel_forward_view_batch(one_vals, one_idx, view_params, **args)) + assert widths == [2] # the body never saw one pixel + for power in (1, 2): + wrapped = padded_back(sino_t, one_idx, view_params, + coeff_power=power, **args) + plain = _parallel_back_view_batch(sino_t, one_idx, view_params, + coeff_power=power, **args) + assert wrapped.shape == plain.shape + assert torch.equal(wrapped, plain), power + assert widths == [2, 2, 2] + + all_idx = torch.as_tensor(idx, dtype=torch.int64) + padded_fwd(torch.as_tensor(vals), all_idx, view_params, **args) + padded_back(sino_t, all_idx, view_params, **args) + assert widths[-2:] == [len(idx), len(idx)] + + # And the driver wraps what it compiles, only that: a hand-written kernel + # body comes back from maybe_compile as itself, cannot be miscompiled, and + # must keep its identity and its cost attribute. + pf = m.projector_functions + raw_fwd, raw_back = m._view_batch_bodies() + for bound, raw in ((pf._fwd_body_per_dev[0], raw_fwd), + (pf._back_body_per_dev[0], raw_back)): + assert bound.__name__.startswith('padded_') == (bound is not raw) + diff --git a/tests/test_widening_floors.py b/tests/test_widening_floors.py index f025e60..7d00ca8 100644 --- a/tests/test_widening_floors.py +++ b/tests/test_widening_floors.py @@ -1,4 +1,5 @@ -"""The widening speed floors, their invariants, and their staleness report. +"""The widening speed floors, their invariants, their staleness report, and +the refresh tool's report of which geometries still need measuring. The floors are a MEASUREMENT of where each device count starts paying for itself, and a measurement is only as good as the code it was taken against. @@ -20,11 +21,35 @@ The selection RULE these numbers feed is tested in test_device_policy.py. """ +import importlib.util +import os import warnings +import pytest + from mbirtorch import _widening_floors as wf REFRESH = 'python dev_scripts/refresh_widening_floors.py' +REFRESH_PATH = os.path.join( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))), + 'dev_scripts', 'refresh_widening_floors.py') + + +@pytest.fixture(scope='module') +def refresh_tool(): + """The refresh script, loaded from its path. + + dev_scripts is not an installed package, so the tool is loaded by file + rather than imported by name. Only its reporting helpers are exercised + here; nothing in this file measures anything or starts a subprocess. + """ + if not os.path.exists(REFRESH_PATH): + pytest.skip('dev_scripts/refresh_widening_floors.py is not present') + spec = importlib.util.spec_from_file_location('refresh_widening_floors', + REFRESH_PATH) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module def install_a_table_with_a_sentinel(monkeypatch): @@ -333,6 +358,94 @@ def test_sinogram_elements_is_the_product_of_the_shape(): assert wf.sinogram_elements((1024, 1008, 992)) == 1_023_934_464 +# ── the refresh tool's "needs measurement" report ──────────────────────────── +def test_the_refresh_tool_reports_the_geometries_that_take_the_fallback( + refresh_tool): + """The one tool whose job is to say "this geometry needs measurement" + must not be silent about the geometries that actually need it. + + A class that declares no floor family is governed by the DEFAULT_FAMILY + floors, which were measured on a different geometry. That is the state + every newly ported geometry arrives in, so it is reported under the None + key rather than skipped for having nothing declared. + """ + import mbirtorch + + missing = refresh_tool.unmeasured_families() + assert None in missing, ( + 'the classes that declare no floor family are the ones taking the ' + 'substituted floors, and they are what this report is for') + undeclared = missing[None] + assert 'TranslationModel' in undeclared + assert 'MultiAxisParallelModel' in undeclared + + # Every reported class really does inherit the base value rather than + # naming a family of its own, so the report matches the code it describes. + for name in undeclared: + assert getattr(mbirtorch, name)._floor_family is None + + +def test_the_report_covers_every_geometry_that_reaches_the_device_decision( + refresh_tool): + """The report is scoped to classes a floor can actually govern. + + A floor is consulted when a model chooses its own device count, which + happens on the shared reconstruction path. The base class is not a + geometry and QGGMRFDenoiser refuses recon, so neither can reach that + decision and neither is work to measure. An exported alias is the same + class object as the class it aliases, so it must not be counted twice. + """ + reported = {name for names in refresh_tool.unmeasured_families().values() + for name in names} + assert 'TomographyModel' not in reported + assert 'QGGMRFDenoiser' not in reported + # MultiAxisParallelBeamModel is an alias of MultiAxisParallelModel; the + # class is reported once, under its own name. + assert 'MultiAxisParallelBeamModel' not in reported + # The two measured families are governed by their own rows, so they are + # not outstanding work. + assert 'ParallelBeamModel' not in reported + assert 'ConeBeamModel' not in reported + + +def test_a_declared_family_with_no_rows_is_still_reported_under_its_name( + refresh_tool, monkeypatch): + """The other way to arrive unmeasured: name a family the table has never + heard of. Widening the report to undeclared classes must not drop the + case it already handled, so both keys are exercised here.""" + monkeypatch.setattr(refresh_tool.wf, 'FLOORS', + {key: value for key, value in wf.FLOORS.items() + if key[0] != 'parallel'}) + + missing = refresh_tool.unmeasured_families() + assert missing.get('parallel') == ['ParallelBeamModel'] + assert 'TranslationModel' in missing[None] + + +def test_the_printed_report_names_the_class_and_the_floors_it_borrows( + refresh_tool, capsys): + """Reading the report has to be enough: it names which class is + unmeasured and which family's floors are standing in for it, so nobody + has to go read the fallback rule to find out what is governing.""" + refresh_tool.print_plan(refresh_tool.build_plan(smoke=True), smoke=True) + + printed = capsys.readouterr().out + assert 'NEEDS MEASUREMENT' in printed + assert 'TranslationModel' in printed + assert 'MultiAxisParallelModel' in printed + assert wf.DEFAULT_FAMILY in printed + + +def test_the_refresh_tool_refuses_to_measure_a_family_it_cannot_build( + refresh_tool): + """The report above invites someone to declare a new floor family. The + builder must then refuse the family it has no geometry for: falling + through to parallel beam would time parallel beam and record the numbers + under the new family's name.""" + with pytest.raises(ValueError, match='cannot build a model for floor family'): + refresh_tool._build_model('translation', (8, 12, 16), 'cpu') + + # ── the env knob ───────────────────────────────────────────────────────────── def test_the_guard_is_on_by_default_and_off_only_when_asked(monkeypatch): monkeypatch.delenv(wf.GUARD_ENV_VAR, raising=False)